diff --git a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java index a5c4c0f86e77..d46d3822f827 100644 --- a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java +++ b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java @@ -34,12 +34,17 @@ import com.google.auth.oauth2.EnvironmentProvider; import com.google.auth.oauth2.PropertyProvider; import com.google.common.base.Strings; +import com.google.common.io.BaseEncoding; import java.io.File; import java.io.FileInputStream; import java.io.IOException; import java.io.InputStream; +import java.nio.file.Files; +import java.nio.file.Paths; +import java.security.MessageDigest; import java.util.Locale; import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; /** * Utility class for mTLS related operations. @@ -57,6 +62,168 @@ private MtlsUtils() { // Prevent instantiation for Utility class } + /** + * Returns if mutual TLS client certificate should be used. Returns true if valid workload + * certificates are configured or if GOOGLE_API_USE_CLIENT_CERTIFICATE is explicitly set to true + * (e.g. for Enterprise Certificate Proxy or custom MtlsProviders), unless explicitly disabled via + * GOOGLE_API_USE_CLIENT_CERTIFICATE=false. + */ + public static boolean useMtlsClientCertificate( + EnvironmentProvider envProvider, PropertyProvider propProvider) { + String useClientCertificate = envProvider.getEnv("GOOGLE_API_USE_CLIENT_CERTIFICATE"); + if ("false".equalsIgnoreCase(useClientCertificate)) { + return false; + } + if (getWorkloadCertPath(envProvider, propProvider) != null) { + return true; + } + return "true".equalsIgnoreCase(useClientCertificate); + } + + /** + * Resolves and returns the path to the mutual TLS client certificate, or null if none should be + * used. + * + *

Possible outcomes: + * + *

    + *
  1. Non-null {@link String} (Valid happy path): A valid workload certificate + * configuration was found and both the certificate and private key files exist and are + * readable. + *
  2. {@link IllegalStateException} (Invalid state - fail closed): An explicit {@code + * GOOGLE_API_CERTIFICATE_CONFIG} path or an existing default well-known certificate + * configuration file is missing, unreadable, malformed, or references missing/unreadable + * certificate or private key files. This is treated as an unrecoverable misconfiguration. + *
  3. {@code null} (Safe fallback / fail open): Client certificates are explicitly + * disabled via {@code GOOGLE_API_USE_CLIENT_CERTIFICATE=false}, no explicit configuration + * is set and the default well-known configuration file does not exist on disk, or the + * configuration specifies an non-workload source (e.g., ECP/PKCS11 without a {@code + * workload} section). Callers can proceed without workload certificate file polling. + *
+ */ + public static @Nullable String getWorkloadCertPath( + EnvironmentProvider envProvider, PropertyProvider propProvider) { + String useClientCertificate = envProvider.getEnv("GOOGLE_API_USE_CLIENT_CERTIFICATE"); + if ("false".equalsIgnoreCase(useClientCertificate)) { + return null; + } + + String explicitConfigPath = envProvider.getEnv(CERTIFICATE_CONFIGURATION_ENV_VARIABLE); + + // 1. Explicit Configuration Path (Fail Closed) + if (!Strings.isNullOrEmpty(explicitConfigPath)) { + File configFile = new File(explicitConfigPath); + if (!configFile.exists()) { + throw new IllegalStateException( + "Certificate configuration file specified via GOOGLE_API_CERTIFICATE_CONFIG at '" + + explicitConfigPath + + "' does not exist."); + } + if (!configFile.isFile() || !configFile.canRead()) { + throw new IllegalStateException( + "Failed to read certificate configuration file specified via" + + " GOOGLE_API_CERTIFICATE_CONFIG at '" + + explicitConfigPath + + "'."); + } + WorkloadCertificateConfiguration config; + try { + config = getWorkloadCertificateConfiguration(envProvider, propProvider, explicitConfigPath); + } catch (CertificateSourceUnavailableException e) { + // ECP / PKCS11 configuration without workload section; safe fallback + return null; + } catch (Exception e) { + throw new IllegalStateException( + "Certificate configuration file specified via GOOGLE_API_CERTIFICATE_CONFIG at '" + + explicitConfigPath + + "' is malformed: " + + e.getMessage(), + e); + } + checkCertAndKeyFilesReadable(config, explicitConfigPath, false); + return config.getCertPath(); + } + + // 2. Implicit / Default gcloud Configuration Path + File defaultConfigFile = null; + try { + defaultConfigFile = getWellKnownCertificateConfigFile(envProvider, propProvider); + } catch (IOException e) { + // APPDATA missing on Windows, etc. Safe fallback. + } + if (defaultConfigFile != null && defaultConfigFile.exists()) { + if (!defaultConfigFile.isFile() || !defaultConfigFile.canRead()) { + throw new IllegalStateException( + "Default certificate configuration file at '" + + defaultConfigFile.getAbsolutePath() + + "' exists but could not be read."); + } + WorkloadCertificateConfiguration config = null; + try { + config = getWorkloadCertificateConfiguration(envProvider, propProvider, null); + } catch (CertificateSourceUnavailableException e) { + // ECP-only configuration without workload section; safe fallback + } catch (Exception e) { + throw new IllegalStateException( + "Default certificate configuration file at '" + + defaultConfigFile.getAbsolutePath() + + "' is malformed: " + + e.getMessage(), + e); + } + if (config != null) { + checkCertAndKeyFilesReadable(config, defaultConfigFile.getAbsolutePath(), true); + return config.getCertPath(); + } + } + + return null; + } + + private static void checkCertAndKeyFilesReadable( + WorkloadCertificateConfiguration config, String configPath, boolean isDefaultConfig) { + File certFile = new File(config.getCertPath()); + File keyFile = new File(config.getPrivateKeyPath()); + if (!certFile.isFile() || !certFile.canRead() || !keyFile.isFile() || !keyFile.canRead()) { + String sourcePrefix = + isDefaultConfig + ? "referenced by default configuration '" + : "referenced by configuration '"; + throw new IllegalStateException( + "Failed to read certificate/key file at '" + + config.getCertPath() + + "' or '" + + config.getPrivateKeyPath() + + "' " + + sourcePrefix + + configPath + + "'."); + } + } + + /** + * Computes the lower-case SHA-256 hex fingerprint of the certificate file at {@code certPath}. + * + *

Unlike {@link #getWorkloadCertPath}, which validates configuration at channel initialization + * and fails closed on errors, this method is called dynamically at runtime during active RPCs to + * detect certificate rotations on disk. External certificate rotators may temporarily delete, + * truncate, or rewrite the certificate file mid-RPC. Returning {@code null} on read/digest + * exceptions (which callers normalize to {@code ""}) allows runtime refresh checks to ignore + * transient mid-write states and keep the active healthy channel without failing in-flight RPCs. + */ + public static @Nullable String getCertificateFingerprint(@Nullable String certPath) { + if (certPath == null) { + return null; + } + try { + byte[] certBytes = Files.readAllBytes(Paths.get(certPath)); + byte[] digest = MessageDigest.getInstance("SHA-256").digest(certBytes); + return BaseEncoding.base16().lowerCase().encode(digest); + } catch (Exception e) { + return null; + } + } + /** * Returns the path to the client certificate file specified by the loaded workload certificate * configuration. @@ -65,14 +232,17 @@ private MtlsUtils() { * @throws IOException if the certificate configuration cannot be found or loaded. */ public static String getCertificatePath( - EnvironmentProvider envProvider, PropertyProvider propProvider, String certConfigPathOverride) + EnvironmentProvider envProvider, + PropertyProvider propProvider, + @Nullable String certConfigPathOverride) throws IOException { String certPath = getWorkloadCertificateConfiguration(envProvider, propProvider, certConfigPathOverride) .getCertPath(); if (Strings.isNullOrEmpty(certPath)) { throw new CertificateSourceUnavailableException( - "Certificate configuration loaded successfully, but does not contain a 'certificate_file' path."); + "Certificate configuration loaded successfully, but does not contain a" + + " 'cert_configs.workload.cert_path' path."); } return certPath; } @@ -92,7 +262,9 @@ public static String getCertificatePath( * @throws IOException if the configuration file cannot be found, read, or parsed */ static WorkloadCertificateConfiguration getWorkloadCertificateConfiguration( - EnvironmentProvider envProvider, PropertyProvider propProvider, String certConfigPathOverride) + EnvironmentProvider envProvider, + PropertyProvider propProvider, + @Nullable String certConfigPathOverride) throws IOException { File certConfig; if (certConfigPathOverride != null) { @@ -106,7 +278,7 @@ static WorkloadCertificateConfiguration getWorkloadCertificateConfiguration( } } - if (!certConfig.isFile()) { + if (!certConfig.isFile() || !certConfig.canRead()) { throw new CertificateSourceUnavailableException( "Certificate configuration file does not exist or is not a file: " + certConfig.getAbsolutePath()); diff --git a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java index f3fdf05a4c32..d5fdc840d804 100644 --- a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java +++ b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java @@ -101,6 +101,21 @@ public String getProperty(String name, String def) { () -> MtlsUtils.getCertificatePath(envProvider, propProvider, configFile.toString())); } + @Test + void getCertificatePath_ecpOnlyConfig_throwsCertificateSourceUnavailableException() + throws IOException { + Path configFile = tempDir.resolve("ecp_config.json"); + Files.write( + configFile, "{\"cert_configs\":{\"enterprise_certificates\":{\"libs\":[]}}}".getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = (name, def) -> def; + + assertThrows( + CertificateSourceUnavailableException.class, + () -> MtlsUtils.getCertificatePath(envProvider, propProvider, configFile.toString())); + } + @Test void getWorkloadCertificateConfiguration_overridePath() throws IOException { Path configFile = tempDir.resolve("custom_config.json"); @@ -243,4 +258,484 @@ public String getProperty(String name, String def) { assertEquals("APPDATA environment variable is not set on Windows.", exception.getMessage()); } + + @Test + void + useMtlsClientCertificate_trueWithNoCertsOnDisk_returnsTrueWhileWorkloadCertPathReturnsNull() { + EnvironmentProvider envProvider = + name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "true" : null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void useMtlsClientCertificate_trueWithEcpOnlyConfig_returnsTrueAndWorkloadCertPathReturnsNull() + throws IOException { + Path configFile = tempDir.resolve("ecp_config.json"); + Files.write(configFile, "{\"cert_configs\":{\"enterprise_certificates\":{}}}".getBytes()); + + EnvironmentProvider envProvider = + name -> { + if ("GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name)) return "true"; + if ("GOOGLE_API_CERTIFICATE_CONFIG".equals(name)) return configFile.toString(); + return null; + }; + PropertyProvider propProvider = (name, def) -> def; + + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void useMtlsClientCertificate_false_returnsFalse() { + EnvironmentProvider envProvider = + name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "false" : null; + PropertyProvider propProvider = (name, def) -> def; + + assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void useMtlsClientCertificate_falseEvenWhenWorkloadCertsExist_returnsFalse() throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(certFile, "dummy cert".getBytes()); + Files.write(keyFile, "dummy key".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certFile.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> { + if ("GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name)) return "false"; + if ("GOOGLE_API_CERTIFICATE_CONFIG".equals(name)) return configFile.toString(); + return null; + }; + PropertyProvider propProvider = (name, def) -> def; + + assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void useMtlsClientCertificate_unsetWithNoCertsOnDisk_returnsFalse() { + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + // --- Explicit GOOGLE_API_CERTIFICATE_CONFIG Tests (Fail Closed) --- + + @Test + void getWorkloadCertPath_explicitConfigMissing_throwsIllegalStateException() { + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? "/nonexistent/config.json" : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "specified via GOOGLE_API_CERTIFICATE_CONFIG at '/nonexistent/config.json' does not" + + " exist")); + } + + @Test + void getWorkloadCertPath_explicitConfigIsDirectory_throwsIllegalStateException() + throws IOException { + Path configDir = tempDir.resolve("config_dir"); + Files.createDirectory(configDir); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configDir.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Failed to read certificate configuration file specified via" + + " GOOGLE_API_CERTIFICATE_CONFIG")); + } + + @Test + void getWorkloadCertPath_explicitConfigUnreadable_throwsIllegalStateException() + throws IOException { + Path configFile = tempDir.resolve("unreadable_config.json"); + Files.write(configFile, "{}".getBytes()); + File file = configFile.toFile(); + if (file.setReadable(false)) { + try { + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Failed to read certificate configuration file specified via" + + " GOOGLE_API_CERTIFICATE_CONFIG")); + } finally { + file.setReadable(true); + } + } + } + + @Test + void getWorkloadCertPath_explicitConfigMalformedJson_throwsIllegalStateException() + throws IOException { + Path configFile = tempDir.resolve("malformed.json"); + Files.write(configFile, "{ invalid json".getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "specified via GOOGLE_API_CERTIFICATE_CONFIG at '" + + configFile.toString() + + "' is malformed")); + } + + @Test + void getWorkloadCertPath_explicitConfigOnlyEcp_returnsNullSafely() throws IOException { + Path configFile = tempDir.resolve("ecp_config.json"); + Files.write(configFile, "{\"cert_configs\":{\"enterprise_certificates\":{}}}".getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void getWorkloadCertPath_explicitConfigCertFileMissing_throwsIllegalStateException() + throws IOException { + Path keyFile = tempDir.resolve("key.pem"); + Files.write(keyFile, "dummy key".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"/nonexistent/cert.pem\",\"key_path\":\"%s\"}}}", + keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + assertTrue( + exception + .getMessage() + .contains("referenced by configuration '" + configFile.toString() + "'")); + } + + @Test + void getWorkloadCertPath_explicitConfigCertFileIsDirectory_throwsIllegalStateException() + throws IOException { + Path certDir = tempDir.resolve("cert_dir"); + Files.createDirectory(certDir); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(keyFile, "dummy key".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certDir.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + } + + @Test + void getWorkloadCertPath_explicitConfigKeyFileMissing_throwsIllegalStateException() + throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Files.write(certFile, "dummy cert".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"/nonexistent/key.pem\"}}}", + certFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + assertTrue( + exception + .getMessage() + .contains("referenced by configuration '" + configFile.toString() + "'")); + } + + @Test + void getWorkloadCertPath_explicitConfigValid_returnsCertPath() throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(certFile, "dummy cert".getBytes()); + Files.write(keyFile, "dummy key".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certFile.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertEquals(certFile.toString(), MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + // --- Implicit / Default gcloud Config Tests --- + + @Test + void getWorkloadCertPath_defaultConfigMissing_returnsNullSafely() { + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void getWorkloadCertPath_defaultConfigIsDirectory_throwsIllegalStateException() + throws IOException { + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + Files.createDirectory(defaultConfigFile); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Default certificate configuration file at '" + + defaultConfigFile.toFile().getAbsolutePath() + + "' exists but could not be read")); + } + + @Test + void getWorkloadCertPath_defaultConfigMalformedJson_throwsIllegalStateException() + throws IOException { + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + Files.write(defaultConfigFile, "{ malformed json".getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Default certificate configuration file at '" + + defaultConfigFile.toFile().getAbsolutePath() + + "' is malformed")); + } + + @Test + void getWorkloadCertPath_defaultConfigOnlyEcp_returnsNullSafely() throws IOException { + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + Files.write( + defaultConfigFile, + "{\"cert_configs\":{\"enterprise_certificates\":{\"libs\":[]}}}".getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void getWorkloadCertPath_defaultConfigCertFileMissing_throwsIllegalStateException() + throws IOException { + Path keyFile = tempDir.resolve("key.pem"); + Files.write(keyFile, "dummy key".getBytes()); + + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"/nonexistent/cert.pem\",\"key_path\":\"%s\"}}}", + keyFile.toString().replace("\\", "\\\\")); + Files.write(defaultConfigFile, configJson.getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + assertTrue( + exception + .getMessage() + .contains( + "referenced by default configuration '" + + defaultConfigFile.toFile().getAbsolutePath() + + "'")); + } + + @Test + void getWorkloadCertPath_defaultConfigValid_returnsCertPath() throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(certFile, "dummy cert".getBytes()); + Files.write(keyFile, "dummy key".getBytes()); + + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certFile.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); + Files.write(defaultConfigFile, configJson.getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertEquals(certFile.toString(), MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + // --- General Helpers & Stubs Tests --- + + @Test + void getCertificateFingerprint_validFile_returnsSha256() throws IOException { + Path file = tempDir.resolve("test.crt"); + Files.write(file, "hello world".getBytes()); + + String fingerprint = MtlsUtils.getCertificateFingerprint(file.toString()); + assertNotNull(fingerprint); + assertEquals(64, fingerprint.length()); // SHA-256 hex string length + assertEquals("b94d27b9934d3e08a52e52d7da7dabfac484efe37a5380ee9088f7ace2efcde9", fingerprint); + } + + @Test + void getCertificateFingerprint_emptyFile_returnsValidSha256() throws IOException { + Path emptyFile = tempDir.resolve("empty.crt"); + Files.write(emptyFile, new byte[0]); + + String fingerprint = MtlsUtils.getCertificateFingerprint(emptyFile.toString()); + assertEquals("e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", fingerprint); + } + + @Test + void getCertificateFingerprint_invalidOrNull_returnsNull() { + assertNull(MtlsUtils.getCertificateFingerprint(null)); + assertNull(MtlsUtils.getCertificateFingerprint("/nonexistent/file.crt")); + assertNull(MtlsUtils.getCertificateFingerprint(tempDir.toString())); // Directory + } } diff --git a/sdk-platform-java/gax-java/gax-grpc/pom.xml b/sdk-platform-java/gax-java/gax-grpc/pom.xml index d717ac0944b5..f5cc9b3798a3 100644 --- a/sdk-platform-java/gax-java/gax-grpc/pom.xml +++ b/sdk-platform-java/gax-java/gax-grpc/pom.xml @@ -162,7 +162,7 @@ maven-surefire-plugin - !InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSetViaEnv_warns,!InstantiatingGrpcChannelProviderTest#canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue,InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfigWrongCredential + !InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSetViaEnv_warns,!InstantiatingGrpcChannelProviderTest#canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index fab73a55dccf..bc6149e74a58 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -31,6 +31,7 @@ import com.google.api.core.InternalApi; import com.google.api.gax.core.FixedExecutorProvider; +import com.google.api.gax.rpc.mtls.CertificateRotationTracker; import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Preconditions; import com.google.common.collect.ImmutableList; @@ -53,9 +54,11 @@ import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; import java.util.concurrent.atomic.AtomicReference; import java.util.logging.Level; import java.util.logging.Logger; +import javax.annotation.concurrent.GuardedBy; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -72,18 +75,28 @@ @NullMarked class ChannelPool extends ManagedChannel { static final String CHANNEL_POOL_CONSECUTIVE_RESIZING_WARNING = - "The gRPC ChannelPool used in the client has been flagged to be repeatedly resizing (5+ times). See https://github.com/googleapis/google-cloud-java/blob/main/docs/grpc_channel_pool_guide.md for more information about this behavior."; + "The gRPC ChannelPool used in the client has been flagged to be repeatedly resizing (5+" + + " times). See" + + " https://github.com/googleapis/google-cloud-java/blob/main/docs/grpc_channel_pool_guide.md" + + " for more information about this behavior."; @VisibleForTesting static final Logger LOG = Logger.getLogger(ChannelPool.class.getName()); private static final java.time.Duration REFRESH_PERIOD = java.time.Duration.ofMinutes(50); private final ChannelPoolSettings settings; private final ChannelFactory channelFactory; private final FixedExecutorProvider backgroundExecutorProvider; + private final String workloadCertPath; private @Nullable ScheduledFuture refreshFuture = null; private @Nullable ScheduledFuture resizeFuture = null; + private final CertificateRotationTracker rotationTracker; private final Object entryWriteLock = new Object(); + + @GuardedBy("entryWriteLock") + private boolean isShutdown = false; + + private final AtomicLong generation = new AtomicLong(0); @VisibleForTesting final AtomicReference> entries = new AtomicReference<>(); private final AtomicInteger indexTicker = new AtomicInteger(); private final String authority; @@ -100,14 +113,15 @@ class ChannelPool extends ManagedChannel { static ChannelPool create( ChannelPoolSettings settings, ChannelFactory channelFactory, - @Nullable ScheduledExecutorService backgroundExecutor) + @Nullable ScheduledExecutorService backgroundExecutor, + @Nullable String workloadCertPath) throws IOException { FixedExecutorProvider executorProvider = backgroundExecutor == null ? FixedExecutorProvider.create(Executors.newSingleThreadScheduledExecutor(), true) : FixedExecutorProvider.create(backgroundExecutor, false); - return new ChannelPool(settings, channelFactory, executorProvider); + return new ChannelPool(settings, channelFactory, executorProvider, workloadCertPath); } /** @@ -121,11 +135,14 @@ static ChannelPool create( ChannelPool( ChannelPoolSettings settings, ChannelFactory channelFactory, - FixedExecutorProvider executorProvider) + FixedExecutorProvider executorProvider, + @Nullable String workloadCertPath) throws IOException { this.settings = settings; this.channelFactory = channelFactory; this.backgroundExecutorProvider = executorProvider; + this.workloadCertPath = workloadCertPath; + this.rotationTracker = new CertificateRotationTracker(workloadCertPath); ImmutableList.Builder initialListBuilder = ImmutableList.builder(); @@ -187,7 +204,8 @@ public ManagedChannel shutdown() { // Resize and refresh tasks can block on channel priming. We don't need // to wait for the channels to be ready since we're shutting down the - // pool. Allowing interrupt to speed it up. + // pool. Allowing interrupt to speed it up. This is done before acquiring + // entryWriteLock, which a running refresh or resize holds. if (resizeFuture != null) { resizeFuture.cancel(true); } @@ -195,9 +213,12 @@ public ManagedChannel shutdown() { refreshFuture.cancel(true); } - List localEntries = entries.get(); - for (Entry entry : localEntries) { - entry.channel.shutdown(); + synchronized (entryWriteLock) { + isShutdown = true; + List localEntries = entries.get(); + for (Entry entry : localEntries) { + entry.channel.shutdown(); + } } if (backgroundExecutorProvider.shouldAutoClose()) { @@ -210,6 +231,11 @@ public ManagedChannel shutdown() { /** {@inheritDoc} */ @Override public boolean isShutdown() { + synchronized (entryWriteLock) { + if (isShutdown) { + return true; + } + } List localEntries = entries.get(); for (Entry entry : localEntries) { if (!entry.channel.isShutdown()) { @@ -236,6 +262,8 @@ public boolean isTerminated() { public ManagedChannel shutdownNow() { LOG.fine("Initiating immediate shutdown due to explicit request"); + // Cancel before acquiring entryWriteLock, which a running refresh or resize holds, so that + // they are interrupted instead of delaying shutdown. if (resizeFuture != null) { resizeFuture.cancel(true); } @@ -243,9 +271,12 @@ public ManagedChannel shutdownNow() { refreshFuture.cancel(true); } - List localEntries = entries.get(); - for (Entry entry : localEntries) { - entry.channel.shutdownNow(); + synchronized (entryWriteLock) { + isShutdown = true; + List localEntries = entries.get(); + for (Entry entry : localEntries) { + entry.channel.shutdownNow(); + } } if (backgroundExecutorProvider.shouldAutoClose()) { @@ -411,7 +442,7 @@ private void expand(int desiredSize) { for (int i = 0; i < desiredSize - localEntries.size(); i++) { try { newEntries.add(new Entry(channelFactory.createSingleChannel())); - } catch (IOException e) { + } catch (Exception e) { LOG.log(Level.WARNING, "Failed to add channel", e); } } @@ -419,23 +450,81 @@ private void expand(int desiredSize) { entries.set(newEntries.build()); } + /** + * Periodically refreshes all channels when {@link + * ChannelPoolSettings#isPreemptiveRefreshEnabled()} is enabled (to mitigate hourly GFE + * disconnects). This applies to all channels even when {@code workloadCertPath == null}. If + * {@code workloadCertPath} is configured, the refresh is skipped while the certificate file is + * unreadable or mid-write on disk. + * + *

If the certificate on disk differs from the active one, this refresh is treated like a + * rotation refresh: channels that fail to be recreated are dropped (and the pool refilled), so + * the switch is completed and the generation advanced only once every channel left in the pool + * uses the new certificate. Otherwise a channel that fails to be recreated keeps its slot, and + * the generation is not advanced, so that a periodic refresh without a rotation does not make + * in-flight {@code UNAUTHENTICATED} failures eligible for a rotation retry. + */ private void refreshSafely() { try { - refresh(); + // See refresh(): a generation change while waiting for the lock means a concurrent refresh + // just recreated every channel on a new certificate. + long generationBeforeLock = generation.get(); + synchronized (entryWriteLock) { + if (workloadCertPath == null) { + refreshAll(); + return; + } + if (generation.get() != generationBeforeLock) { + LOG.fine( + "Skipping pre-emptive channel refresh: channels were just recreated by a concurrent" + + " refresh"); + return; + } + String currentDiskFingerprint = rotationTracker.readDiskFingerprint(); + if (currentDiskFingerprint.isEmpty()) { + // The configured certificate could not be read, which normally means it is being + // rewritten during a rotation. Skip this refresh rather than recreate channels from a + // partially written certificate; the next periodic or reactive refresh recreates them. + LOG.fine( + "Skipping pre-emptive channel refresh: the workload certificate could not be read;" + + " channels will be recreated on the next refresh"); + return; + } + boolean certChanged = !rotationTracker.isAlreadyActive(currentDiskFingerprint); + if (refreshAll(/* dropUnrefreshedChannels= */ certChanged) && certChanged) { + completeCertificateSwitch(currentDiskFingerprint); + } + } } catch (Exception e) { - LOG.log(Level.WARNING, "Failed to pre-emptively refresh channnels", e); + LOG.log(Level.WARNING, "Failed to pre-emptively refresh channels", e); } } + @VisibleForTesting + void invalidateDiskFingerprintCache() { + rotationTracker.invalidateCache(); + } + + boolean shouldRefresh() { + return rotationTracker.shouldRefresh(); + } + /** * Replace all of the channels in the channel pool with fresh ones. This is meant to mitigate the - * hourly GFE disconnects by giving clients the ability to prime the channel on reconnect. + * hourly GFE disconnects by giving clients the ability to prime the channel on reconnect, and to + * pick up a rotated mTLS workload certificate. * - *

This is done on a best effort basis. If the replacement channel fails to construct, the old - * channel will continue to be used. + *

This is done on a best effort basis. When no workload certificate is configured, a channel + * whose replacement fails to construct continues to be used. When a workload certificate is + * configured, a channel whose replacement fails to construct is dropped so that no traffic is + * routed to a channel using the old certificate; if no replacement can be constructed at all, the + * pool is left unchanged. */ @InternalApi("Visible for testing") void refresh() { + // A generation change while waiting for the lock means a concurrent refresh already switched + // every channel to a new certificate, so there is no need to read the certificate again. + long generationBeforeLock = generation.get(); // Note: synchronization is necessary in case refresh is called concurrently: // - thread1 fails to replace a single entry // - thread2 succeeds replacing an entry @@ -443,28 +532,179 @@ void refresh() { // - then thread2 will shut down channel that thread1 will put back into circulation (after it // replaces the list) synchronized (entryWriteLock) { + if (isShutdown) { + return; + } + if (workloadCertPath == null) { + refreshAll(); + return; + } + if (generation.get() != generationBeforeLock) { + LOG.fine( + "Channel pool was already refreshed by a concurrent thread, skipping duplicate" + + " refresh"); + return; + } + String currentDiskFingerprint = rotationTracker.readDiskFingerprint(); + if (currentDiskFingerprint.isEmpty()) { + return; + } + + // Double-check fingerprint inside the lock + if (rotationTracker.isAlreadyActive(currentDiskFingerprint)) { + LOG.fine( + "Channel pool was already refreshed by a concurrent thread, skipping duplicate" + + " refresh"); + return; + } + + // Drop any channel that fails to refresh so that no traffic is routed to the old certificate. + if (refreshAll(/* dropUnrefreshedChannels= */ true)) { + completeCertificateSwitch(currentDiskFingerprint); + } + } + } + + /** + * Records that every channel in the pool now uses the certificate with the given fingerprint. + * Must be called while holding {@code entryWriteLock}, after the channels have been swapped. + * + *

The generation is incremented before the fingerprint is marked active, so that a concurrent + * failing RPC that no longer sees a pending rotation ({@link #shouldRefresh()} is {@code false}) + * is guaranteed to see the new generation and be retried on the new channels. + */ + private void completeCertificateSwitch(String newFingerprint) { + generation.incrementAndGet(); + rotationTracker.markRefreshed(newFingerprint); + LOG.fine("Channel pool switched to certificate with fingerprint: " + newFingerprint); + } + + @InternalApi("Visible for testing") + @Nullable String getWorkloadCertPath() { + return workloadCertPath; + } + + @InternalApi("Visible for testing") + boolean refreshAll() { + return refreshAll(/* dropUnrefreshedChannels= */ false); + } + + /** + * Replaces the channels in the pool with freshly created ones. + * + * @param dropUnrefreshedChannels if {@code false}, a channel that fails to be recreated keeps its + * slot in the pool. If {@code true} (used for certificate rotation), only newly created + * channels are kept so no traffic is routed to a channel using the old certificate, and the + * pool is then refilled asynchronously to its size before the refresh. + * @return if {@code dropUnrefreshedChannels} is {@code false}, whether every channel was + * recreated; otherwise, whether at least one channel was recreated (i.e. every channel left + * in the pool was newly created) + */ + @InternalApi("Visible for testing") + boolean refreshAll(boolean dropUnrefreshedChannels) { + synchronized (entryWriteLock) { + if (isShutdown) { + return false; + } LOG.fine("Refreshing all channels"); - ArrayList newEntries = new ArrayList<>(entries.get()); + // All writes to entries happen under entryWriteLock, so this snapshot is the current pool. + ImmutableList currentEntries = entries.get(); + List keptEntries = new ArrayList<>(currentEntries.size()); + List retiredEntries = new ArrayList<>(currentEntries.size()); + List createdEntries = new ArrayList<>(currentEntries.size()); + boolean allCreated = !currentEntries.isEmpty(); - for (int i = 0; i < newEntries.size(); i++) { - try { - newEntries.set(i, new Entry(channelFactory.createSingleChannel())); - } catch (IOException e) { - LOG.log(Level.WARNING, "Failed to refresh channel, leaving old channel", e); + try { + for (Entry oldEntry : currentEntries) { + try { + Entry newEntry = new Entry(channelFactory.createSingleChannel()); + createdEntries.add(newEntry); + keptEntries.add(newEntry); + retiredEntries.add(oldEntry); + } catch (Exception e) { + allCreated = false; + if (dropUnrefreshedChannels) { + retiredEntries.add(oldEntry); + } else { + keptEntries.add(oldEntry); + } + LOG.log( + Level.WARNING, + dropUnrefreshedChannels + ? "Failed to refresh channel, dropping old channel" + : "Failed to refresh channel, leaving old channel", + e); + } } - } - ImmutableList replacedEntries = entries.getAndSet(ImmutableList.copyOf(newEntries)); + if (createdEntries.isEmpty()) { + return false; + } + + entries.set(ImmutableList.copyOf(keptEntries)); + createdEntries.clear(); // Ownership transferred to pool + + // Shutdown the channels that were cycled out. + retiredEntries.forEach(Entry::requestShutdown); - // Shutdown the channels that were cycled out. - for (Entry e : replacedEntries) { - if (!newEntries.contains(e)) { + // Restore the pool to its pre-refresh size right away, rather than leave a dynamically + // sized + // pool to grow back over several resize() runs while traffic queues on fewer channels. + if (dropUnrefreshedChannels && !allCreated) { + scheduleRefill(currentEntries.size()); + } + return dropUnrefreshedChannels || allCreated; + } finally { + // If an Error aborted before the swap, shut down newly created channels so they don't leak + for (Entry e : createdEntries) { e.requestShutdown(); } } } } + /** + * Schedules a one-shot task that restores the pool to {@code targetSize} channels after a + * certificate rotation refresh dropped channels that failed to refresh. + */ + private void scheduleRefill(int targetSize) { + try { + backgroundExecutorProvider.getExecutor().execute(() -> refillSafely(targetSize)); + } catch (RuntimeException e) { + LOG.log(Level.WARNING, "Failed to schedule channel pool refill", e); + } + } + + private void refillSafely(int targetSize) { + try { + synchronized (entryWriteLock) { + if (isShutdown) { + return; + } + if (entries.get().size() < targetSize) { + expand(targetSize); + } + } + } catch (Exception e) { + LOG.log(Level.WARNING, "Failed to refill channel pool", e); + } + } + + /** + * Returns the current channel pool generation counter. + * + *

The generation is a monotonically increasing counter incremented each time the pool switches + * every channel to a new mTLS certificate (a reactive refresh after a rotation, or a periodic + * refresh that picks up a rotated certificate). Periodic refreshes that do not change the + * certificate do not increment it. Retry loops ({@code AttemptCallable} and {@code + * ServerStreamingAttemptCallable}) snapshot the generation before starting an RPC attempt and + * compare it after an {@code UNAUTHENTICATED} failure to determine whether the pool rotated to a + * new certificate during or after the attempt. + */ + long getGeneration() { + return generation.get(); + } + /** * Get and retain a Channel Entry. The returned Entry will have its rpc count incremented, * preventing it from getting recycled. @@ -616,17 +856,30 @@ public ClientCall newCall( MethodDescriptor methodDescriptor, CallOptions callOptions) { Entry entry = getRetainedEntry(affinity); - - return new ReleasingClientCall<>(entry.channel.newCall(methodDescriptor, callOptions), entry); + try { + return new ReleasingClientCall<>( + entry.channel.newCall(methodDescriptor, callOptions), entry); + } catch (Throwable t) { + entry.release(); + throw t; + } } } - /** ClientCall wrapper that makes sure to decrement the outstanding RPC count on completion. */ + /** + * ClientCall wrapper that makes sure to decrement the outstanding RPC count on completion. + * + *

Contract: Exactly one call to {@link #start(Listener, Metadata)} is required to balance + * reference counts. Early cancellation before {@code start()} is recorded and safely decrements + * the reference count when {@code start()} is subsequently invoked. + */ static class ReleasingClientCall extends SimpleForwardingClientCall { - private @Nullable CancellationException cancellationException; + private final Object callLock = new Object(); + private volatile @Nullable CancellationException cancellationException; final Entry entry; private final AtomicBoolean wasClosed = new AtomicBoolean(); private final AtomicBoolean wasReleased = new AtomicBoolean(); + private final AtomicBoolean wasStarted = new AtomicBoolean(); public ReleasingClientCall(ClientCall delegate, Entry entry) { super(delegate); @@ -635,51 +888,81 @@ public ReleasingClientCall(ClientCall delegate, Entry entry) { @Override public void start(Listener responseListener, Metadata headers) { - if (cancellationException != null) { - throw new IllegalStateException("Call is already cancelled", cancellationException); - } - try { - super.start( - new SimpleForwardingClientCallListener(responseListener) { - @Override - public void onClose(Status status, Metadata trailers) { - if (!wasClosed.compareAndSet(false, true)) { - LOG.log( - Level.WARNING, - "Call is being closed more than once. Please make sure that onClose() is not being manually called."); - return; - } - try { - super.onClose(status, trailers); - } finally { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); - } else { + synchronized (callLock) { + if (!wasStarted.compareAndSet(false, true)) { + throw new IllegalStateException("Call is already started"); + } + if (cancellationException != null) { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } + throw new IllegalStateException("Call is already cancelled", cancellationException); + } + try { + super.start( + new SimpleForwardingClientCallListener(responseListener) { + @Override + public void onClose(Status status, Metadata trailers) { + if (!wasClosed.compareAndSet(false, true)) { LOG.log( Level.WARNING, - "Entry was released before the call is closed. This may be due to an exception on start of the call."); + "Call is being closed more than once. Please make sure that onClose() is" + + " not being manually called."); + return; + } + try { + super.onClose(status, trailers); + } finally { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } else { + LOG.log( + Level.WARNING, + "Entry was released before the call is closed. This may be due to an" + + " exception on start of the call."); + } } } - } - }, - headers); - } catch (Exception e) { - // In case start failed, make sure to release - if (wasReleased.compareAndSet(false, true)) { - entry.release(); - } else { - LOG.log( - Level.WARNING, - "The entry is already released. This indicates that onClose() has already been called previously"); + }, + headers); + } catch (Throwable t) { + // In case start failed, make sure to release + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } else { + LOG.log( + Level.WARNING, + "The entry is already released. This indicates that onClose() has already been" + + " called previously"); + } + throw t; } - throw e; } } @Override public void cancel(@Nullable String message, @Nullable Throwable cause) { - this.cancellationException = new CancellationException(message); - super.cancel(message, cause); + boolean releaseImmediately = false; + try { + synchronized (callLock) { + this.cancellationException = new CancellationException(message); + if (!wasStarted.get()) { + releaseImmediately = true; + } + if (delegate() != null) { + super.cancel(message, cause); + } + } + } catch (Throwable t) { + if (!wasStarted.get()) { + releaseImmediately = true; + } + throw t; + } finally { + if (releaseImmediately && wasReleased.compareAndSet(false, true)) { + entry.release(); + } + } } } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java index 23f56c5f8951..fab685328931 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java @@ -99,6 +99,7 @@ public final class GrpcCallContext implements ApiCallContext { private final ApiCallContextOptions options; private final EndpointContext endpointContext; private final boolean isDirectPath; + @Nullable private final TransportChannel transportChannel; /** Returns an empty instance with a null channel and default {@link CallOptions}. */ public static GrpcCallContext createDefault() { @@ -115,7 +116,8 @@ public static GrpcCallContext createDefault() { null, null, null, - false); + false, + null); } /** Returns an instance with the given channel and {@link CallOptions}. */ @@ -133,7 +135,8 @@ public static GrpcCallContext of(Channel channel, CallOptions callOptions) { null, null, null, - false); + false, + null); } private GrpcCallContext( @@ -149,7 +152,8 @@ private GrpcCallContext( @Nullable RetrySettings retrySettings, @Nullable Set retryableCodes, @Nullable EndpointContext endpointContext, - boolean isDirectPath) { + boolean isDirectPath, + @Nullable TransportChannel transportChannel) { this.channel = channel; this.credentials = credentials; Preconditions.checkNotNull(callOptions); @@ -169,6 +173,7 @@ private GrpcCallContext( this.endpointContext = endpointContext == null ? EndpointContext.getDefaultInstance() : endpointContext; this.isDirectPath = isDirectPath; + this.transportChannel = transportChannel; } /** @@ -210,7 +215,13 @@ public GrpcCallContext withCredentials(Credentials newCredentials) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); + } + + @Override + public TransportChannel getTransportChannel() { + return transportChannel; } @Override @@ -234,7 +245,8 @@ public GrpcCallContext withTransportChannel(TransportChannel inputChannel) { retrySettings, retryableCodes, endpointContext, - transportChannel.isDirectPath()); + transportChannel.isDirectPath(), + inputChannel); } @Override @@ -253,7 +265,8 @@ public GrpcCallContext withEndpointContext(EndpointContext endpointContext) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } /** This method is obsolete. Use {@link #withTimeoutDuration(java.time.Duration)} instead. */ @@ -288,7 +301,8 @@ public GrpcCallContext withTimeoutDuration(java.time.@Nullable Duration timeout) retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -334,7 +348,8 @@ public GrpcCallContext withStreamWaitTimeoutDuration( retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } /** @@ -369,7 +384,8 @@ public GrpcCallContext withStreamIdleTimeoutDuration( retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @BetaApi("The surface for channel affinity is not stable yet and may change in the future.") @@ -387,7 +403,8 @@ public GrpcCallContext withChannelAffinity(@Nullable Integer affinity) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @BetaApi("The surface for extra headers is not stable yet and may change in the future.") @@ -409,7 +426,8 @@ public GrpcCallContext withExtraHeaders(Map> extraHeaders) retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -432,7 +450,8 @@ public GrpcCallContext withRetrySettings(RetrySettings retrySettings) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -455,7 +474,8 @@ public GrpcCallContext withRetryableCodes(Set retryableCodes) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -542,6 +562,12 @@ public ApiCallContext merge(ApiCallContext inputCallContext) { newCallOptions = newCallOptions.withOption(TRACER_KEY, newTracer); } + TransportChannel newTransportChannel = grpcCallContext.transportChannel; + if (newTransportChannel == null + && (grpcCallContext.channel == null || grpcCallContext.channel.equals(channel))) { + newTransportChannel = transportChannel; + } + // The EndpointContext is not updated as there should be no reason for a user // to update this. return new GrpcCallContext( @@ -557,7 +583,8 @@ public ApiCallContext merge(ApiCallContext inputCallContext) { newRetrySettings, newRetryableCodes, endpointContext, - newIsDirectPath); + newIsDirectPath, + newTransportChannel); } /** The {@link Channel} set on this context. */ @@ -635,7 +662,8 @@ public GrpcCallContext withChannel(@Nullable Channel newChannel) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + (newChannel != null && newChannel.equals(channel)) ? transportChannel : null); } /** Returns a new instance with the call options set to the given call options. */ @@ -653,7 +681,8 @@ public GrpcCallContext withCallOptions(CallOptions newCallOptions) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } public GrpcCallContext withRequestParamsDynamicHeaderOption(String requestParams) { @@ -698,7 +727,8 @@ public GrpcCallContext withOption(Key key, T value) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } /** {@inheritDoc} */ @@ -759,7 +789,8 @@ public int hashCode() { options, retrySettings, retryableCodes, - endpointContext); + endpointContext, + transportChannel); } @Override @@ -783,7 +814,8 @@ public boolean equals(@Nullable Object o) { && Objects.equals(options, that.options) && Objects.equals(retrySettings, that.retrySettings) && Objects.equals(retryableCodes, that.retryableCodes) - && Objects.equals(endpointContext, that.endpointContext); + && Objects.equals(endpointContext, that.endpointContext) + && Objects.equals(transportChannel, that.transportChannel); } Metadata getMetadata() { diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java index e0a520facb17..63180a8149f6 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java @@ -68,6 +68,32 @@ public Channel getChannel() { return getManagedChannel(); } + @Override + public void refresh() { + Channel channel = getChannel(); + if (channel instanceof ChannelPool) { + ((ChannelPool) channel).refresh(); + } + } + + @Override + public boolean shouldRefresh() { + Channel channel = getChannel(); + if (channel instanceof ChannelPool) { + return ((ChannelPool) channel).shouldRefresh(); + } + return false; + } + + @Override + public long getGeneration() { + Channel channel = getChannel(); + if (channel instanceof ChannelPool) { + return ((ChannelPool) channel).getGeneration(); + } + return 0; + } + @Override public void shutdown() { getManagedChannel().shutdown(); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java index ac42396f006a..67fc57b075d4 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java @@ -401,12 +401,21 @@ public TransportChannel getTransportChannel() throws IOException { } private TransportChannel createChannel() throws IOException { + // Only track the workload certificate for rotation when the pool's channels actually present + // it, mirroring the credential selection in createSingleChannel(): DirectPath channels use + // GoogleDefaultChannelCredentials (ALTS) rather than the client certificate, and without an + // mtlsProvider there is no client certificate KeyStore (S2A or plain TLS is used instead). + String workloadCertPath = + !this.canUseDirectPath() && mtlsProvider != null + ? certificateBasedAccess.getWorkloadCertPath() + : null; return GrpcTransportChannel.newBuilder() .setManagedChannel( ChannelPool.create( channelPoolSettings, InstantiatingGrpcChannelProvider.this::createSingleChannel, - backgroundExecutor)) + backgroundExecutor, + workloadCertPath)) .setDirectPath(this.canUseDirectPath()) .build(); } @@ -465,8 +474,9 @@ private void logDirectPathMisconfig() { level, "Env var " + DIRECT_PATH_ENV_ENABLE_XDS - + " was found and set to TRUE, but DirectPath was not enabled for this client. If this is intended for " - + "this client, please note that this is a misconfiguration and set the attemptDirectPath option as well."); + + " was found and set to TRUE, but DirectPath was not enabled for this client. If" + + " this is intended for this client, please note that this is a misconfiguration" + + " and set the attemptDirectPath option as well."); } // Case 2: Direct Path xDS was enabled via Builder. Direct Path Traffic Director must be set // (enabled with `setAttemptDirectPath(true)`) along with xDS. @@ -474,7 +484,9 @@ private void logDirectPathMisconfig() { else if (isDirectPathXdsEnabledViaBuilderOption()) { LOG.log( level, - "DirectPath is misconfigured. The DirectPath XDS option was set, but the attemptDirectPath option was not. Please set both the attemptDirectPath and attemptDirectPathXds options."); + "DirectPath is misconfigured. The DirectPath XDS option was set, but the" + + " attemptDirectPath option was not. Please set both the attemptDirectPath and" + + " attemptDirectPathXds options."); } } else { // Case 3: credential is not correctly set @@ -666,7 +678,8 @@ ChannelCredentials createS2ASecuredChannelCredentials() { // Fallback to plaintext connection to S2A. LOG.log( Level.INFO, - "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not return a mtls address to reach S2A."); + "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not" + + " return a mtls address to reach S2A."); s2aChannelCredentials = createPlaintextToS2AChannelCredentials(plaintextAddress); return s2aChannelCredentials; } @@ -685,7 +698,9 @@ ChannelCredentials createS2ASecuredChannelCredentials() { // Fallback to plaintext-to-S2A connection on error. LOG.log( Level.WARNING, - "Cannot establish an mTLS connection to S2A due to error creating MTLS to MDS TlsChannelCredentials credentials, falling back to plaintext connection to S2A: " + "Cannot establish an mTLS connection to S2A due to error creating MTLS to MDS" + + " TlsChannelCredentials credentials, falling back to plaintext connection to" + + " S2A: " + ignore.getMessage()); s2aChannelCredentials = createPlaintextToS2AChannelCredentials(plaintextAddress); return s2aChannelCredentials; @@ -755,6 +770,8 @@ public ManagedChannelBuilder createChannelBuilder() throws IOException { if (channelCredentials != null) { // Create the channel using channel credentials created via DCA. builder = Grpc.newChannelBuilder(endpoint, channelCredentials); + } else if (mtlsProvider != null && certificateBasedAccess.useMtlsClientCertificate()) { + throw new IOException("Failed to initialize mTLS channel credentials"); } else { // Could not create channel credentials via DCA. In accordance with // https://google.aip.dev/auth/4115, if credentials not available through @@ -1403,7 +1420,8 @@ public InstantiatingGrpcChannelProvider build() { "DefaultMtlsProviderFactory encountered unexpected IOException: " + e.getMessage()); LOG.log( Level.WARNING, - "mTLS configuration was detected on the device, but mTLS failed to initialize. Falling back to non-mTLS channel."); + "mTLS configuration was detected on the device, but mTLS failed to initialize." + + " Falling back to non-mTLS channel."); } } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 5bfdc7754759..ec2bacff40e6 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -47,6 +47,7 @@ import com.google.api.gax.rpc.StreamController; import com.google.api.gax.rpc.UnaryCallSettings; import com.google.api.gax.rpc.UnaryCallable; +import com.google.api.gax.rpc.mtls.CertificateRotationTracker; import com.google.api.gax.util.FakeLogHandler; import com.google.auth.Credentials; import com.google.common.collect.ImmutableList; @@ -65,12 +66,16 @@ import java.util.Arrays; import java.util.List; import java.util.concurrent.CancellationException; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.logging.Level; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; @@ -81,13 +86,18 @@ class ChannelPoolTest { private static final int DEFAULT_AWAIT_TERMINATION_SEC = 10; private ChannelPool pool; + private java.nio.file.Path tempCert; @AfterEach - void cleanup() throws InterruptedException { + void cleanup() throws InterruptedException, IOException { if (pool != null) { pool.shutdown(); pool.awaitTermination(DEFAULT_AWAIT_TERMINATION_SEC, TimeUnit.SECONDS); } + if (tempCert != null) { + java.nio.file.Files.deleteIfExists(tempCert); + tempCert = null; + } } @Test @@ -101,6 +111,7 @@ void testAuthority() throws IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(2), new FakeChannelFactory(Arrays.asList(sub1, sub2)), + null, null); assertThat(pool.authority()).isEqualTo("myAuth"); } @@ -117,6 +128,7 @@ void testRoundRobin() throws IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(channels.size()), new FakeChannelFactory(channels), + null, null); verifyTargetChannel(pool, channels, sub1); @@ -195,6 +207,7 @@ void ensureEvenDistribution() throws InterruptedException, IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(numChannels), new FakeChannelFactory(Arrays.asList(channels)), + null, null); int numThreads = 20; @@ -233,6 +246,7 @@ void channelPrimerShouldCallPoolConstruction() throws IOException { .setPreemptiveRefreshEnabled(true) .build(), new FakeChannelFactory(Arrays.asList(channel1, channel2), mockChannelPrimer), + null, null); Mockito.verify(mockChannelPrimer, Mockito.times(2)) .primeChannel(Mockito.any(ManagedChannel.class)); @@ -273,7 +287,8 @@ void channelPrimerIsCalledPeriodically() throws IOException { .setPreemptiveRefreshEnabled(true) .build(), channelFactory, - provider); + provider, + null); // 1 call during the creation Mockito.verify(mockChannelPrimer, Mockito.times(1)) .primeChannel(Mockito.any(ManagedChannel.class)); @@ -297,7 +312,7 @@ void callShouldCompleteAfterCreation() throws IOException { ManagedChannel replacementChannel = mock(ManagedChannel.class); FakeChannelFactory channelFactory = new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); - pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); // create a mock call when new call comes to the underlying channel MockClientCall mockClientCall = new MockClientCall<>(1, Status.OK); @@ -322,7 +337,7 @@ void callShouldCompleteAfterCreation() throws IOException { ClientCall call = pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); - pool.refresh(); + pool.refreshAll(); // shutdown is not called because there is still an outstanding call, even if it hasn't started Mockito.verify(underlyingChannel, Mockito.after(200).never()).shutdown(); @@ -346,7 +361,7 @@ void callShouldCompleteAfterStarted() throws IOException { FakeChannelFactory channelFactory = new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); - pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); // create a mock call when new call comes to the underlying channel MockClientCall mockClientCall = new MockClientCall<>(1, Status.OK); @@ -373,7 +388,7 @@ void callShouldCompleteAfterStarted() throws IOException { // start clientCall call.start(listener, new Metadata()); - pool.refresh(); + pool.refreshAll(); // shutdown is not called because there is still an outstanding call Mockito.verify(underlyingChannel, Mockito.after(200).never()).shutdown(); @@ -391,7 +406,7 @@ void channelShouldShutdown() throws IOException { FakeChannelFactory channelFactory = new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); - pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); // create a mock call when new call comes to the underlying channel MockClientCall mockClientCall = new MockClientCall<>(1, Status.OK); @@ -422,11 +437,978 @@ void channelShouldShutdown() throws IOException { call.sendMessage("message"); // shutdown is not called because it has not been shutdown yet Mockito.verify(underlyingChannel, Mockito.after(200).never()).shutdown(); - pool.refresh(); + pool.refreshAll(); // shutdown is called because the outstanding call has completed Mockito.verify(underlyingChannel, Mockito.atLeastOnce()).shutdown(); } + @Test + void testCancelBeforeStartReleasesChannelEntry() throws IOException { + ManagedChannel underlyingChannel = mock(ManagedChannel.class); + ManagedChannel replacementChannel = mock(ManagedChannel.class); + FakeChannelFactory channelFactory = + new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + ClientCall call = + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + + pool.refreshAll(); + Mockito.verify(underlyingChannel, Mockito.never()).shutdown(); + + call.cancel("Cancelled early", null); + Mockito.verify(underlyingChannel, Mockito.times(1)).shutdown(); + } + + @Test + void channelReactiveMTlsRefresh_swapsChannelsOnlyWhenCertChanges() + throws IOException, InterruptedException { + ManagedChannel underlyingChannel1 = Mockito.mock(ManagedChannel.class); + ManagedChannel underlyingChannel2 = Mockito.mock(ManagedChannel.class); + + FakeChannelFactory channelFactory = + new FakeChannelFactory(ImmutableList.of(underlyingChannel1, underlyingChannel2)); + + // Create a temp file to act as the cert + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + + java.nio.file.Path clientCert = + java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); + java.nio.file.Files.copy( + clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + ChannelPoolSettings channelPoolSettings = + ChannelPoolSettings.builder().setInitialChannelCount(1).build(); + + pool = ChannelPool.create(channelPoolSettings, channelFactory, null, tempCert.toString()); + + // Initially uses channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel1, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Try a reactive refresh *without* changing the cert content (should no-op) + pool.refresh(); + + // Verify it's STILL channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel1, Mockito.times(2)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // The ChannelPool caches fingerprints for 1000ms, wait for it to expire + pool.invalidateDiskFingerprintCache(); + + java.nio.file.Path rootCert = + java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); + java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + // Try a reactive refresh *with* a changed cert content (should swap channels) + pool.refresh(); + + // Verify it is NOW channel2 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel2, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + } + + @Test + void channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsRetry() + throws IOException { + ManagedChannel channel1 = Mockito.mock(ManagedChannel.class); + ManagedChannel channel2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + + // Initial creation returns channel1, refresh attempt 1 throws IOException, refresh attempt 2 + // returns channel2 + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(channel1) + .thenThrow(new IOException("Transient channel creation error")) + .thenReturn(channel2); + + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + java.nio.file.Path clientCert = + java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); + java.nio.file.Files.copy( + clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + ChannelPoolSettings channelPoolSettings = + ChannelPoolSettings.builder().setInitialChannelCount(1).build(); + + pool = ChannelPool.create(channelPoolSettings, channelFactory, null, tempCert.toString()); + + // Initially uses channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(channel1, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Rotate cert on disk + pool.invalidateDiskFingerprintCache(); + java.nio.file.Path rootCert = + java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); + java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + // Refresh attempt 1: createSingleChannel throws IOException. + // Refresh should fail to replace channel and MUST NOT record the new cert fingerprint as + // active. + pool.refresh(); + + // Verify still channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(channel1, Mockito.times(2)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Refresh attempt 2: with the same cert file on disk (cache expired), channelFactory now + // succeeds. + // If the fingerprint had been mutated on the failed attempt, this call would be skipped as a + // duplicate! + pool.invalidateDiskFingerprintCache(); + pool.refresh(); + + // Verify it has now swapped to channel2! + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(channel2, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + } + + private void writeCert(String resourceName) throws IOException { + if (tempCert == null) { + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + } + java.nio.file.Files.copy( + java.nio.file.Paths.get("src", "test", "resources", resourceName), + tempCert, + java.nio.file.StandardCopyOption.REPLACE_EXISTING); + } + + /** Creates an mTLS pool backed by {@code executor} and then rotates the certificate on disk. */ + private ChannelPool createMtlsPoolAndRotateCert( + ChannelPoolSettings settings, + ChannelFactory channelFactory, + ScheduledExecutorService executor) + throws IOException { + writeCert("client_cert.pem"); + pool = + new ChannelPool( + settings, channelFactory, FixedExecutorProvider.create(executor), tempCert.toString()); + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + assertThat(pool.shouldRefresh()).isTrue(); + return pool; + } + + private static ScheduledExecutorService mockExecutor() { + return Mockito.mock( + ScheduledExecutorService.class, Mockito.withSettings().withoutAnnotations()); + } + + private static ChannelFactory mockChannelFactory() { + return Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + } + + @Test + void channelReactiveMTlsRefresh_partialFailureInStaticPool_dropsUnrefreshedChannelsAndRefills() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ManagedChannel refilled = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(rotated1) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(refilled); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + long genBefore = pool.getGeneration(); + + pool.refresh(); + + // Only the newly created channel is kept; both channels on the old certificate are retired. + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + Mockito.verify(rotated1, Mockito.never()).shutdown(); + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(rotated1) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + assertThat(pool.getGeneration()).isEqualTo(genBefore + 1); + // Every channel left in the pool uses the new certificate, so it is recorded as active. + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + + // A one-time refill restores the configured size. + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(channelFactory, Mockito.times(5)).createSingleChannel(); + + // Running the refill again on a full pool creates no channels. + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(channelFactory, Mockito.times(5)).createSingleChannel(); + Mockito.verify(refilled, Mockito.never()).shutdown(); + Mockito.verify(executor, Mockito.times(1)).execute(Mockito.any(Runnable.class)); + } + + @Test + void channelReactiveMTlsRefresh_fullSuccess_doesNotScheduleRefill() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2, rotated1, rotated2); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + + pool.refresh(); + + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + } + + @Test + void refreshAll_partialFailureWithoutRotation_keepsOldChannelAndDoesNotScheduleRefill() + throws IOException { + // Non-rotation refreshes (e.g. the preemptive refresh) keep per-slot fallback even for mTLS. + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel refreshed1 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(refreshed1) + .thenThrow(new IOException("Transient failure on second sub-channel")); + + writeCert("client_cert.pem"); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(2), + channelFactory, + FixedExecutorProvider.create(executor), + tempCert.toString()); + long genBefore = pool.getGeneration(); + List before = pool.entries.get(); + + assertThat(pool.refreshAll()).isFalse(); + + assertThat(pool.entries.get()).hasSize(2); + // Each channel keeps its slot, so affinity-to-index mapping is unchanged. + assertThat(pool.entries.get().get(0)).isNotSameInstanceAs(before.get(0)); + assertThat(pool.entries.get().get(1)).isSameInstanceAs(before.get(1)); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2, Mockito.never()).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(genBefore); + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + } + + @Test + void channelReactiveMTlsRefresh_allChannelsFail_leavesPoolUnchangedAndDoesNotRefill() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenThrow(new IOException("Failure on first sub-channel")) + .thenThrow(new IOException("Failure on second sub-channel")); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + long genBefore = pool.getGeneration(); + List before = pool.entries.get(); + + pool.refresh(); + + // Nothing could be rebuilt, so the pool keeps its old channels and the rotation stays pending. + assertThat(pool.entries.get()).containsExactlyElementsIn(before).inOrder(); + Mockito.verify(initial1, Mockito.never()).shutdown(); + Mockito.verify(initial2, Mockito.never()).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(genBefore); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isTrue(); + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + } + + @Test + void channelReactiveMTlsRefresh_partialFailureInDynamicPool_refillsToPreRefreshSize() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ManagedChannel refilled = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(rotated1) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(refilled); + + createMtlsPoolAndRotateCert( + ChannelPoolSettings.builder() + .setInitialChannelCount(2) + .setMinChannelCount(2) + .setMaxChannelCount(4) + .setMinRpcsPerChannel(1) + .setMaxRpcsPerChannel(2) + .build(), + channelFactory, + executor); + + pool.refresh(); + + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + + // Dynamic pools are refilled right away too, rather than waiting for resize(). + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(refilled, Mockito.never()).shutdown(); + } + + @Test + void channelReactiveMTlsRefresh_partialFailureAfterResize_refillsToPreRefreshSize() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn( + // initial pool of 4 + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + // rotation refresh: 1 of the 2 remaining channels is recreated + Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(Mockito.mock(ManagedChannel.class)); + + createMtlsPoolAndRotateCert( + ChannelPoolSettings.builder() + .setInitialChannelCount(4) + .setMinChannelCount(2) + .setMaxChannelCount(6) + .setMinRpcsPerChannel(1) + .setMaxRpcsPerChannel(2) + .build(), + channelFactory, + executor); + // With no load, resize() shrinks the pool below its initial channel count. + pool.resize(); + assertThat(pool.entries.get()).hasSize(2); + + pool.refresh(); + assertThat(pool.entries.get()).hasSize(1); + + // The refill targets the size before the refresh, not the initial channel count. + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(channelFactory, Mockito.times(7)).createSingleChannel(); + } + + @Test + void refill_afterShutdown_createsNoChannels() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn( + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(Mockito.mock(ManagedChannel.class)); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + pool.refresh(); + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + + pool.shutdown(); + refillTask.getValue().run(); + + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(channelFactory, Mockito.times(4)).createSingleChannel(); + } + + @Test + void refill_channelCreationFailure_isLoggedAndDoesNotThrow() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn( + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenThrow(new IOException("Checked failure during refill")) + .thenThrow(new RuntimeException("Unchecked failure during refill")) + .thenReturn(Mockito.mock(ManagedChannel.class)); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + pool.refresh(); + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + + // Checked and unchecked failures are both handled by expand(); the pool stays usable at its + // reduced size. + FakeLogHandler logHandler = new FakeLogHandler(); + ChannelPool.LOG.addHandler(logHandler); + try { + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(1); + + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(1); + } finally { + ChannelPool.LOG.removeHandler(logHandler); + } + assertThat(logHandler.getAllMessages()) + .containsExactly("Failed to add channel", "Failed to add channel"); + + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + } + + @Test + void refill_failureAfterPartialProgress_keepsCreatedChannels() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel refilled = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn( + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenThrow(new IOException("Transient failure on third sub-channel")) + .thenReturn(refilled) + .thenThrow(new RuntimeException("Unchecked failure during refill")); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(3), channelFactory, executor); + pool.refresh(); + assertThat(pool.entries.get()).hasSize(1); + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + + refillTask.getValue().run(); + + // The channel created before the failure is added to the pool rather than orphaned. + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(refilled, Mockito.never()).shutdown(); + Mockito.verify(channelFactory, Mockito.times(8)).createSingleChannel(); + } + + @Test + void channelReactiveMTlsRefresh_refillRejectedByExecutor_stillCompletesRefresh() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + Mockito.doThrow(new RejectedExecutionException("Executor shut down")) + .when(executor) + .execute(Mockito.any(Runnable.class)); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2, Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + long genBefore = pool.getGeneration(); + + pool.refresh(); + + Mockito.verify(executor).execute(Mockito.any(Runnable.class)); + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(genBefore + 1); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + } + + @Test + void refreshAll_runtimeExceptionOrError_doesNotLeakCreatedChannels() throws IOException { + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel createdBeforeRuntimeEx = Mockito.mock(ManagedChannel.class); + ManagedChannel createdBeforeError = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(createdBeforeRuntimeEx) + .thenThrow(new RuntimeException("Unchecked runtime exception")) + .thenReturn(createdBeforeError) + .thenThrow(new AssertionError("Simulated Error during refresh")); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(2), channelFactory, null, null); + + // Case 1: RuntimeException on channel 1 after creating channel 0 -> caught as Exception, + // partial progress committed + boolean allCreated = pool.refreshAll(); + assertThat(allCreated).isFalse(); + Mockito.verify(initial1).shutdown(); + + // Case 2: Error on channel 1 after creating channel 0 -> aborts, finally block must shut down + // createdBeforeError + org.junit.jupiter.api.Assertions.assertThrows(AssertionError.class, () -> pool.refreshAll()); + Mockito.verify(createdBeforeError).shutdown(); + } + + @Test + void refresh_onShutdownPool_noOpsAndCreatesNoChannels() throws IOException { + ManagedChannel channel1 = mock(ManagedChannel.class); + ManagedChannel channel2 = mock(ManagedChannel.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(channel1, channel2); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel(); + + pool.shutdown(); + assertThat(pool.isShutdown()).isTrue(); + + // Invoking refresh or refreshAll on shut down pool must no-op and never create new subchannels + pool.refresh(); + boolean refreshed = pool.refreshAll(); + assertThat(refreshed).isFalse(); + Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel(); + assertThat(pool.isShutdown()).isTrue(); + } + + @Test + void refreshAll_doesNotIncrementGeneration() throws IOException { + ManagedChannel channel1 = mock(ManagedChannel.class); + ManagedChannel channel2 = mock(ManagedChannel.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(channel1, channel2); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + assertThat(pool.getGeneration()).isEqualTo(0); + + // Only a switch to a new certificate advances the generation. + assertThat(pool.refreshAll()).isTrue(); + Mockito.verify(channel1).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(0); + } + + /** + * Creates an mTLS pool with preemptive refresh enabled and returns the scheduled periodic refresh + * task. + */ + private Runnable createPreemptiveRefreshMtlsPool(int size, ChannelFactory channelFactory) + throws IOException { + return createPreemptiveRefreshMtlsPool(size, channelFactory, mockExecutor()); + } + + private Runnable createPreemptiveRefreshMtlsPool( + int size, ChannelFactory channelFactory, ScheduledExecutorService executor) + throws IOException { + List refreshTasks = new ArrayList<>(); + Mockito.doAnswer( + invocation -> { + refreshTasks.add(invocation.getArgument(0)); + return Mockito.mock( + ScheduledFuture.class, Mockito.withSettings().withoutAnnotations()); + }) + .when(executor) + .scheduleAtFixedRate( + Mockito.any(Runnable.class), Mockito.anyLong(), Mockito.anyLong(), Mockito.any()); + writeCert("client_cert.pem"); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(size).toBuilder() + .setPreemptiveRefreshEnabled(true) + .build(), + channelFactory, + FixedExecutorProvider.create(executor), + tempCert.toString()); + assertThat(refreshTasks).hasSize(1); + return refreshTasks.get(0); + } + + private String readCertFingerprint() { + return new CertificateRotationTracker(tempCert.toString()).readDiskFingerprint(); + } + + @Test + void preemptiveRefresh_withoutRotation_doesNotIncrementGeneration() throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel refreshed = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, refreshed); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory); + + FakeLogHandler logHandler = new FakeLogHandler(); + Level originalLevel = ChannelPool.LOG.getLevel(); + ChannelPool.LOG.setLevel(Level.FINE); + ChannelPool.LOG.addHandler(logHandler); + try { + pool.invalidateDiskFingerprintCache(); + preemptiveRefresh.run(); + } finally { + ChannelPool.LOG.removeHandler(logHandler); + ChannelPool.LOG.setLevel(originalLevel); + } + + // The channels are replaced, but the certificate did not change. + Mockito.verify(initial).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(0); + assertThat(logHandler.getAllMessages()).contains("Refreshing all channels"); + assertThat(String.join("\n", logHandler.getAllMessages())) + .doesNotContain("Channel pool switched to certificate"); + } + + @Test + void preemptiveRefresh_pickingUpRotatedCert_incrementsGeneration() throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory); + + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + preemptiveRefresh.run(); + + Mockito.verify(initial).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(1); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + } + + @Test + void preemptiveRefresh_whenCertUnreadable_skipsRefreshAndLogs() throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory); + + // An empty certificate file is what a reader sees while the certificate is being rewritten. + java.nio.file.Files.write(tempCert, new byte[0]); + FakeLogHandler logHandler = new FakeLogHandler(); + Level originalLevel = ChannelPool.LOG.getLevel(); + ChannelPool.LOG.setLevel(Level.FINE); + ChannelPool.LOG.addHandler(logHandler); + try { + pool.invalidateDiskFingerprintCache(); + preemptiveRefresh.run(); + } finally { + ChannelPool.LOG.removeHandler(logHandler); + ChannelPool.LOG.setLevel(originalLevel); + } + + Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel(); + Mockito.verify(initial, Mockito.never()).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(0); + assertThat(String.join("\n", logHandler.getAllMessages())) + .contains("Skipping pre-emptive channel refresh"); + assertThat(logHandler.getAllMessages()).doesNotContain("Refreshing all channels"); + } + + @Test + void preemptiveRefresh_partialFailureOnCertChange_dropsOldChannelsAndCompletesSwitch() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ManagedChannel refilled = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(rotated1) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(refilled); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(2, channelFactory, executor); + + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + preemptiveRefresh.run(); + + // The channel that failed to refresh still uses the old certificate, so it is dropped and + // every channel left in the pool uses the new certificate. + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + assertThat(pool.entries.get()).hasSize(1); + assertThat(pool.getGeneration()).isEqualTo(1); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + + // The pool is refilled to its size before the refresh. + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(refilled, Mockito.never()).shutdown(); + } + + @Test + void preemptiveRefresh_allChannelsFailOnCertChange_leavesRotationPending() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenThrow(new IOException("Transient failure")); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(2, channelFactory, executor); + List entriesBefore = pool.entries.get(); + + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + preemptiveRefresh.run(); + + // No channel could be recreated, so the pool is left unchanged and the rotation stays pending + // for the next periodic or reactive refresh. + assertThat(pool.entries.get()).containsExactlyElementsIn(entriesBefore).inOrder(); + Mockito.verify(initial1, Mockito.never()).shutdown(); + Mockito.verify(initial2, Mockito.never()).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(0); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isTrue(); + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + } + + @Test + void preemptiveRefresh_partialFailureWithoutCertChange_keepsOldChannel() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel refreshed1 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(refreshed1) + .thenThrow(new IOException("Transient failure on second sub-channel")); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(2, channelFactory, executor); + ChannelPool.Entry initialEntry2 = pool.entries.get().get(1); + + pool.invalidateDiskFingerprintCache(); + preemptiveRefresh.run(); + + // The certificate did not change, so the channel that failed to refresh keeps its slot. + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2, Mockito.never()).shutdown(); + assertThat(pool.entries.get()).hasSize(2); + assertThat(pool.entries.get().get(1)).isSameInstanceAs(initialEntry2); + assertThat(pool.getGeneration()).isEqualTo(0); + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + } + + /** Waits until {@code thread} is blocked waiting for a monitor lock. */ + private static void awaitBlocked(Thread thread) throws InterruptedException { + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5); + while (thread.getState() != Thread.State.BLOCKED) { + assertThat(System.nanoTime()).isLessThan(deadline); + Thread.sleep(1); + } + } + + /** + * Starts a reactive refresh that holds the pool lock while it recreates the channel on the + * rotated certificate, runs {@code waiter} on another thread until it blocks on the lock, and + * rotates the certificate again before letting the first refresh complete. A waiter that re-read + * the certificate would see the second rotation and refresh again. + */ + private void runWhileConcurrentSwitchHoldsLock( + CountDownLatch switchStarted, CountDownLatch releaseSwitch, Runnable waiter) + throws Exception { + Thread switchingThread = new Thread(pool::refresh); + Thread waitingThread = new Thread(waiter); + try { + switchingThread.start(); + assertThat(switchStarted.await(5, TimeUnit.SECONDS)).isTrue(); + waitingThread.start(); + awaitBlocked(waitingThread); + writeCert("client_cert.pem"); + } finally { + releaseSwitch.countDown(); + } + switchingThread.join(5000); + waitingThread.join(5000); + assertThat(switchingThread.isAlive()).isFalse(); + assertThat(waitingThread.isAlive()).isFalse(); + } + + @Test + void refresh_concurrentSwitchWhileWaitingForLock_skipsWithoutReadingDisk() throws Exception { + CountDownLatch switchStarted = new CountDownLatch(1); + CountDownLatch releaseSwitch = new CountDownLatch(1); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(Mockito.mock(ManagedChannel.class)) + .thenAnswer( + invocation -> { + switchStarted.countDown(); + releaseSwitch.await(); + return Mockito.mock(ManagedChannel.class); + }) + .thenReturn(Mockito.mock(ManagedChannel.class)); + createMtlsPoolAndRotateCert( + ChannelPoolSettings.staticallySized(1), channelFactory, mockExecutor()); + + runWhileConcurrentSwitchHoldsLock(switchStarted, releaseSwitch, pool::refresh); + + // Only the initial channel and the concurrent switch created channels; the waiter skipped. + Mockito.verify(channelFactory, Mockito.times(2)).createSingleChannel(); + assertThat(pool.getGeneration()).isEqualTo(1); + // The second rotation is still detected for the next UNAUTHENTICATED failure. + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isTrue(); + } + + @Test + void preemptiveRefresh_concurrentSwitchWhileWaitingForLock_skipsRefresh() throws Exception { + CountDownLatch switchStarted = new CountDownLatch(1); + CountDownLatch releaseSwitch = new CountDownLatch(1); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(Mockito.mock(ManagedChannel.class)) + .thenAnswer( + invocation -> { + switchStarted.countDown(); + releaseSwitch.await(); + return Mockito.mock(ManagedChannel.class); + }) + .thenReturn(Mockito.mock(ManagedChannel.class)); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory); + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + + runWhileConcurrentSwitchHoldsLock(switchStarted, releaseSwitch, preemptiveRefresh); + + // The channels were just recreated by the concurrent switch, so the periodic refresh skipped. + Mockito.verify(channelFactory, Mockito.times(2)).createSingleChannel(); + assertThat(pool.getGeneration()).isEqualTo(1); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isTrue(); + } + + @Test + void refresh_whenCertUnchanged_noOpsAndDoesNotIncrementGeneration() throws IOException { + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(Mockito.mock(ManagedChannel.class), Mockito.mock(ManagedChannel.class)); + writeCert("client_cert.pem"); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1), + channelFactory, + FixedExecutorProvider.create(mockExecutor()), + tempCert.toString()); + + pool.invalidateDiskFingerprintCache(); + pool.refresh(); + + assertThat(pool.getGeneration()).isEqualTo(0); + Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel(); + } + + @Test + void refresh_onRotation_logsNewCertificateFingerprint() throws IOException { + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(Mockito.mock(ManagedChannel.class), Mockito.mock(ManagedChannel.class)); + writeCert("client_cert.pem"); + String oldFingerprint = readCertFingerprint(); + createMtlsPoolAndRotateCert( + ChannelPoolSettings.staticallySized(1), channelFactory, mockExecutor()); + String newFingerprint = readCertFingerprint(); + assertThat(newFingerprint).isNotEqualTo(oldFingerprint); + + FakeLogHandler logHandler = new FakeLogHandler(); + Level originalLevel = ChannelPool.LOG.getLevel(); + ChannelPool.LOG.setLevel(Level.FINE); + ChannelPool.LOG.addHandler(logHandler); + try { + pool.refresh(); + } finally { + ChannelPool.LOG.removeHandler(logHandler); + ChannelPool.LOG.setLevel(originalLevel); + } + + assertThat(pool.getGeneration()).isEqualTo(1); + assertThat(logHandler.getAllMessages()) + .contains("Channel pool switched to certificate with fingerprint: " + newFingerprint); + assertThat(String.join("\n", logHandler.getAllMessages())).doesNotContain(oldFingerprint); + } + + @Test + void shutdown_interruptsInProgressRefresh() throws Exception { + CountDownLatch refreshStarted = new CountDownLatch(1); + CountDownLatch release = new CountDownLatch(1); + AtomicBoolean refreshInterrupted = new AtomicBoolean(); + AtomicInteger createdChannels = new AtomicInteger(); + ChannelFactory channelFactory = + () -> { + if (createdChannels.getAndIncrement() > 0) { + // The preemptive refresh blocks while creating its replacement channel. + refreshStarted.countDown(); + try { + release.await(); + } catch (InterruptedException e) { + refreshInterrupted.set(true); + Thread.currentThread().interrupt(); + throw new IOException("Interrupted while creating channel", e); + } + } + return Mockito.mock(ManagedChannel.class); + }; + ScheduledExecutorService realExecutor = Executors.newSingleThreadScheduledExecutor(); + ScheduledExecutorService executor = mockExecutor(); + Mockito.doAnswer( + invocation -> + realExecutor.schedule( + (Runnable) invocation.getArgument(0), 0, TimeUnit.MILLISECONDS)) + .when(executor) + .scheduleAtFixedRate( + Mockito.any(Runnable.class), Mockito.anyLong(), Mockito.anyLong(), Mockito.any()); + try { + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1).toBuilder() + .setPreemptiveRefreshEnabled(true) + .build(), + channelFactory, + FixedExecutorProvider.create(executor), + null); + assertThat(refreshStarted.await(5, TimeUnit.SECONDS)).isTrue(); + + // The refresh holds the pool's write lock; shutdown must interrupt it rather than wait. + Assertions.assertTimeoutPreemptively(java.time.Duration.ofSeconds(5), () -> pool.shutdown()); + + assertThat(refreshInterrupted.get()).isTrue(); + assertThat(pool.isShutdown()).isTrue(); + } finally { + release.countDown(); + realExecutor.shutdownNow(); + } + } + @Test void channelRefreshShouldSwapChannels() throws IOException { ManagedChannel underlyingChannel1 = mock(ManagedChannel.class); @@ -450,7 +1432,8 @@ void channelRefreshShouldSwapChannels() throws IOException { .setPreemptiveRefreshEnabled(true) .build(), channelFactory, - provider); + provider, + null); Mockito.reset(underlyingChannel1); pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); @@ -459,10 +1442,41 @@ void channelRefreshShouldSwapChannels() throws IOException { .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); // swap channel - pool.refresh(); + pool.refreshAll(); + + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + + Mockito.verify(underlyingChannel2, Mockito.only()) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + } + + @Test + void testRefreshWithNullWorkloadCertPathSwapsChannel() throws IOException { + ScheduledExecutorService executor = + Mockito.mock(ScheduledExecutorService.class, Mockito.withSettings().withoutAnnotations()); + FixedExecutorProvider provider = FixedExecutorProvider.create(executor); + ManagedChannel underlyingChannel1 = Mockito.mock(ManagedChannel.class); + ManagedChannel underlyingChannel2 = Mockito.mock(ManagedChannel.class); + FakeChannelFactory channelFactory = + new FakeChannelFactory(ImmutableList.of(underlyingChannel1, underlyingChannel2)); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1).toBuilder() + .setPreemptiveRefreshEnabled(true) + .build(), + channelFactory, + provider, + null); + Mockito.reset(underlyingChannel1); pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel1, Mockito.only()) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Calling refresh() when workloadCertPath is null should fall back to refreshAll() + pool.refresh(); + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); Mockito.verify(underlyingChannel2, Mockito.only()) .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); } @@ -486,7 +1500,8 @@ void channelCountShouldNotChangeWhenOutstandingRpcsAreWithinLimits() throws Exce .setMaxRpcsPerChannel(2) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // Start the minimum number of @@ -553,7 +1568,8 @@ void customResizeDeltaIsRespected() throws Exception { .setMaxResizeDelta(5) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // Add 20 RPCs to push expansion @@ -586,7 +1602,8 @@ void removedIdleChannelsAreShutdown() throws Exception { .setMaxRpcsPerChannel(2) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // With no outstanding RPCs, the pool should shrink @@ -614,7 +1631,8 @@ void removedActiveChannelsAreShutdown() throws Exception { .setMaxRpcsPerChannel(2) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // Start 2 RPCs @@ -652,7 +1670,7 @@ void testReleasingClientCallCancelEarly() throws IOException { Mockito.when(fakeChannel.newCall(Mockito.any(), Mockito.any())).thenReturn(mockClientCall); ChannelPoolSettings channelPoolSettings = ChannelPoolSettings.staticallySized(1); ChannelFactory factory = new FakeChannelFactory(ImmutableList.of(fakeChannel)); - pool = ChannelPool.create(channelPoolSettings, factory, null); + pool = ChannelPool.create(channelPoolSettings, factory, null, null); EndpointContext endpointContext = Mockito.mock(EndpointContext.class, Mockito.withSettings().withoutAnnotations()); @@ -717,7 +1735,8 @@ void repeatedResizingLogsWarningOnExpand() throws Exception { .setMaxChannelCount(10) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(1); FakeLogHandler logHandler = new FakeLogHandler(); @@ -769,7 +1788,8 @@ void repeatedResizingLogsWarningOnShrink() throws Exception { .setMaxChannelCount(10) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(10); FakeLogHandler logHandler = new FakeLogHandler(); @@ -805,7 +1825,7 @@ void testDoubleRelease() throws Exception { ChannelPoolSettings channelPoolSettings = ChannelPoolSettings.staticallySized(1); ChannelFactory factory = new FakeChannelFactory(ImmutableList.of(fakeChannel)); - pool = ChannelPool.create(channelPoolSettings, factory, null); + pool = ChannelPool.create(channelPoolSettings, factory, null, null); EndpointContext endpointContext = Mockito.mock(EndpointContext.class, Mockito.withSettings().withoutAnnotations()); @@ -843,7 +1863,8 @@ void testDoubleRelease() throws Exception { // Ensure that the channel pool properly logged the double call and kept the refCount correct assertThat(logHandler.getAllMessages()) .contains( - "Call is being closed more than once. Please make sure that onClose() is not being manually called."); + "Call is being closed more than once. Please make sure that onClose() is not being" + + " manually called."); assertThat(pool.entries.get()).hasSize(1); ChannelPool.Entry entry = pool.entries.get().get(0); assertThat(entry.outstandingRpcs.get()).isEqualTo(0); @@ -879,7 +1900,8 @@ void minChannelsClampedToMaxChannelCountUnderHighLoad() throws Exception { .setMaxChannelCount(5) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(1); // Add 20 RPCs, which would require 10 channels (20/2) @@ -914,7 +1936,8 @@ void maxChannelsClampedToMinChannelCountUnderLowLoad() throws Exception { .setMaxChannelCount(10) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(5); // With no outstanding RPCs, the pool should want to shrink to 0 @@ -924,4 +1947,278 @@ void maxChannelsClampedToMinChannelCountUnderLowLoad() throws Exception { // Should be clamped to minChannelCount = 3 assertThat(pool.entries.get()).hasSize(3); } + + @Test + void shouldRefresh_doesNotCacheNegativeResultAndDetectsSubsequentRotationImmediately() + throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + java.nio.file.Path clientCert = + java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); + java.nio.file.Files.copy( + clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + pool = + ChannelPool.create( + ChannelPoolSettings.staticallySized(1), channelFactory, null, tempCert.toString()); + + // First check returns false (unchanged disk cert) + assertThat(pool.shouldRefresh()).isFalse(); + + // Immediately rotate cert on disk WITHOUT invalidating the 1-second cache + java.nio.file.Path rootCert = + java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); + java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + // Must immediately detect rotation because negative/unchanged disk checks are not cached for 1s + assertThat(pool.shouldRefresh()).isTrue(); + + // Refresh should update activeCertFingerprint and clear any cached positive check + pool.refresh(); + assertThat(pool.shouldRefresh()).isFalse(); + } + + @Test + void newCall_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() + throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())) + .thenThrow(new LinkageError("Simulated native/JNI linkage error")); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + assertThrows(LinkageError.class, () -> pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT)); + + // Rotating the pool should immediately shut down initial channel because its ref count is 0 + pool.refresh(); + Mockito.verify(initial).shutdown(); + } + + @Test + @SuppressWarnings("unchecked") + void start_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ClientCall mockCall = Mockito.mock(ClientCall.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())).thenReturn((ClientCall) mockCall); + Mockito.doThrow(new AssertionError("Simulated Error in start")) + .when(mockCall) + .start(Mockito.any(), Mockito.any()); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + // Rotate pool while call is retained + pool.refresh(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + // Calling start() throws Error, which must release the retained entry and trigger shutdown + assertThrows( + AssertionError.class, + () -> call.start(new ClientCall.Listener() {}, new io.grpc.Metadata())); + Mockito.verify(initial).shutdown(); + } + + @Test + @SuppressWarnings("unchecked") + void cancel_whenDelegateThrowsException_releasesEntryAndShutsDownRetiredChannel() + throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ClientCall mockCall = Mockito.mock(ClientCall.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())).thenReturn((ClientCall) mockCall); + Mockito.doThrow(new RuntimeException("Simulated cancel exception")) + .when(mockCall) + .cancel(Mockito.any(), Mockito.any()); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + pool.refresh(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + assertThrows(RuntimeException.class, () -> call.cancel("cancelled", null)); + Mockito.verify(initial).shutdown(); + } + + @Test + @SuppressWarnings("unchecked") + void concurrentStartAndCancel_neverLeaksOrDoubleReleasesEntry() throws Exception { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())) + .thenAnswer( + invocation -> + new ClientCall() { + private Listener listener; + private boolean cancelled; + + @Override + public synchronized void start( + Listener responseListener, io.grpc.Metadata headers) { + this.listener = responseListener; + if (cancelled) { + responseListener.onClose(io.grpc.Status.CANCELLED, new io.grpc.Metadata()); + } + } + + @Override + public synchronized void cancel(String message, Throwable cause) { + cancelled = true; + if (listener != null) { + listener.onClose(io.grpc.Status.CANCELLED, new io.grpc.Metadata()); + } + } + + @Override + public void request(int numMessages) {} + + @Override + public void halfClose() {} + + @Override + public void sendMessage(Color message) {} + }); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + int iterations = 100; + java.util.concurrent.ExecutorService executor = + java.util.concurrent.Executors.newFixedThreadPool(2); + try { + for (int i = 0; i < iterations; i++) { + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + java.util.concurrent.CyclicBarrier barrier = new java.util.concurrent.CyclicBarrier(2); + java.util.concurrent.Future f1 = + executor.submit( + () -> { + try { + barrier.await(); + call.start(new ClientCall.Listener() {}, new io.grpc.Metadata()); + } catch (Exception ignored) { + } + }); + java.util.concurrent.Future f2 = + executor.submit( + () -> { + try { + barrier.await(); + call.cancel("cancel", null); + } catch (Exception ignored) { + } + }); + f1.get(5, java.util.concurrent.TimeUnit.SECONDS); + f2.get(5, java.util.concurrent.TimeUnit.SECONDS); + } + } finally { + executor.shutdownNow(); + } + + // Every start/cancel race must leave the entry with no outstanding RPCs: a leak would leave it + // positive and a double release would make it negative. + assertThat(pool.entries.get().get(0).outstandingRpcs.get()).isEqualTo(0); + + // Rotate pool: initial channel must shut down cleanly, proving outstandingRpcs == 0 (no leaks + // or negative counts) + pool.refresh(); + Mockito.verify(initial).shutdown(); + } + + @Test + void cancel_whenStartedAndSuperCancelThrows_doesNotReleasePrematurelyUntilOnClose() + throws Exception { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel replacement = Mockito.mock(ManagedChannel.class); + @SuppressWarnings("unchecked") + ClientCall delegateCall = Mockito.mock(ClientCall.class); + @SuppressWarnings("unchecked") + ArgumentCaptor> listenerCaptor = + ArgumentCaptor.forClass(ClientCall.Listener.class); + Mockito.doThrow(new RuntimeException("cancel failure")) + .when(delegateCall) + .cancel(Mockito.any(), Mockito.any()); + Mockito.when(initial.newCall(Mockito.eq(METHOD_RECOGNIZE), Mockito.any())) + .thenReturn(delegateCall); + + java.util.concurrent.atomic.AtomicInteger createCount = + new java.util.concurrent.atomic.AtomicInteger(0); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1), + () -> createCount.getAndIncrement() == 0 ? initial : replacement, + FixedExecutorProvider.create(Mockito.mock(ScheduledExecutorService.class)), + null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + call.start(new ClientCall.Listener() {}, new Metadata()); + Mockito.verify(delegateCall).start(listenerCaptor.capture(), Mockito.any()); + + assertThrows(RuntimeException.class, () -> call.cancel("abort", null)); + + // Rotate pool while call is still active (onClose hasn't fired yet): + // initial channel must NOT be shut down yet because call is still active + pool.refreshAll(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + // Once onClose fires, entry is released and initial channel shuts down + listenerCaptor.getValue().onClose(Status.CANCELLED, new Metadata()); + Mockito.verify(initial).shutdown(); + } + + @Test + void start_whenCalledTwice_throwsIllegalStateExceptionAndDoesNotReleaseFirstCallEntry() + throws Exception { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel replacement = Mockito.mock(ManagedChannel.class); + @SuppressWarnings("unchecked") + ClientCall delegateCall = Mockito.mock(ClientCall.class); + @SuppressWarnings("unchecked") + ArgumentCaptor> listenerCaptor = + ArgumentCaptor.forClass(ClientCall.Listener.class); + Mockito.when(initial.newCall(Mockito.eq(METHOD_RECOGNIZE), Mockito.any())) + .thenReturn(delegateCall); + + java.util.concurrent.atomic.AtomicInteger createCount = + new java.util.concurrent.atomic.AtomicInteger(0); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1), + () -> createCount.getAndIncrement() == 0 ? initial : replacement, + FixedExecutorProvider.create(Mockito.mock(ScheduledExecutorService.class)), + null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + call.start(new ClientCall.Listener() {}, new Metadata()); + Mockito.verify(delegateCall).start(listenerCaptor.capture(), Mockito.any()); + + // Duplicate start() must throw IllegalStateException without releasing the entry + assertThrows( + IllegalStateException.class, + () -> call.start(new ClientCall.Listener() {}, new Metadata())); + + pool.refreshAll(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + listenerCaptor.getValue().onClose(Status.OK, new Metadata()); + Mockito.verify(initial).shutdown(); + } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java index e20767fdb8ed..55f473407ec8 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java @@ -210,6 +210,22 @@ void testWithLongerTimeout() { .isEqualTo(java.time.Duration.ofSeconds(5)); } + @Test + void testWithNullOrZeroTimeoutClearsExistingTimeout() { + GrpcCallContext ctxWithTimeout = + GrpcCallContext.createDefault().withTimeoutDuration(java.time.Duration.ofSeconds(5)); + + // Sanity check + Truth.assertThat(ctxWithTimeout.getTimeoutDuration()) + .isEqualTo(java.time.Duration.ofSeconds(5)); + + java.time.Duration nullTimeout = null; + Truth.assertThat(ctxWithTimeout.withTimeoutDuration(nullTimeout).getTimeoutDuration()).isNull(); + Truth.assertThat( + ctxWithTimeout.withTimeoutDuration(java.time.Duration.ZERO).getTimeoutDuration()) + .isNull(); + } + @Test void testMergeWithNullTimeout() { java.time.Duration timeout = java.time.Duration.ofSeconds(10); @@ -494,4 +510,64 @@ private static Map> createTestExtraHeaders(String... keyVal } return extraHeaders; } + + @Test + public void testEqualsAndHashCode() { + ManagedChannel managedChannel1 = org.mockito.Mockito.mock(ManagedChannel.class); + ManagedChannel managedChannel2 = org.mockito.Mockito.mock(ManagedChannel.class); + + GrpcTransportChannel transportChannel1 = GrpcTransportChannel.create(managedChannel1); + GrpcTransportChannel transportChannel2 = GrpcTransportChannel.create(managedChannel2); + + GrpcCallContext context1 = + GrpcCallContext.createDefault().withTransportChannel(transportChannel1); + GrpcCallContext context2 = + GrpcCallContext.createDefault().withTransportChannel(transportChannel1); + GrpcCallContext context3 = + GrpcCallContext.createDefault().withTransportChannel(transportChannel2); + + org.junit.jupiter.api.Assertions.assertEquals(context1, context2); + org.junit.jupiter.api.Assertions.assertEquals(context1.hashCode(), context2.hashCode()); + + org.junit.jupiter.api.Assertions.assertNotEquals(context1, context3); + } + + @Test + public void testMergeWithCustomChannelClearsTransportChannel() { + ManagedChannel defaultChannel = org.mockito.Mockito.mock(ManagedChannel.class); + ManagedChannel customChannel = org.mockito.Mockito.mock(ManagedChannel.class); + GrpcTransportChannel transportChannel = GrpcTransportChannel.create(defaultChannel); + + GrpcCallContext baseContext = + GrpcCallContext.createDefault().withTransportChannel(transportChannel); + GrpcCallContext overrideContext = GrpcCallContext.of(customChannel, CallOptions.DEFAULT); + + GrpcCallContext mergedContext = (GrpcCallContext) baseContext.merge(overrideContext); + assertEquals(customChannel, mergedContext.getChannel()); + assertNull(mergedContext.getTransportChannel()); + } + + @Test + public void testWithChannelWithCustomChannelClearsTransportChannel() { + ManagedChannel defaultChannel = org.mockito.Mockito.mock(ManagedChannel.class); + ManagedChannel customChannel = org.mockito.Mockito.mock(ManagedChannel.class); + GrpcTransportChannel transportChannel = GrpcTransportChannel.create(defaultChannel); + + GrpcCallContext baseContext = + GrpcCallContext.createDefault().withTransportChannel(transportChannel); + GrpcCallContext updatedContext = baseContext.withChannel(customChannel); + + assertEquals(customChannel, updatedContext.getChannel()); + assertNull(updatedContext.getTransportChannel()); + + // Clearing channel via withChannel(null) also clears transportChannel + GrpcCallContext nullChannelContext = baseContext.withChannel(null); + assertNull(nullChannelContext.getChannel()); + assertNull(nullChannelContext.getTransportChannel()); + + // Merging a cleared context into defaultContext falls back to defaultContext's transportChannel + GrpcCallContext mergedWithNullChannel = (GrpcCallContext) baseContext.merge(nullChannelContext); + assertEquals(defaultChannel, mergedWithNullChannel.getChannel()); + assertEquals(transportChannel, mergedWithNullChannel.getTransportChannel()); + } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java index 2aa9279e249f..6877eb1dbe1e 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java @@ -125,6 +125,7 @@ void testAffinity() throws IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(2), new FakeChannelFactory(Arrays.asList(channel0, channel1)), + null, null); GrpcCallContext context = defaultCallContext.withChannel(pool); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java index fad4cd468b95..c93db599d575 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java @@ -32,7 +32,6 @@ import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.spy; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; @@ -83,7 +82,7 @@ void testInterceptor_basic() { void testInterceptor_responseListener() { when(channel.newCall(Mockito.>any(), any(CallOptions.class))) .thenReturn(call); - GrpcLoggingInterceptor interceptor = spy(new GrpcLoggingInterceptor()); + GrpcLoggingInterceptor interceptor = new GrpcLoggingInterceptor(); Channel intercepted = ClientInterceptors.intercept(channel, interceptor); @SuppressWarnings("unchecked") ClientCall.Listener listener = mock(ClientCall.Listener.class); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java index be0365866615..d766b288d77d 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java @@ -48,6 +48,7 @@ import com.google.api.gax.rpc.internal.EnvironmentProvider; import com.google.api.gax.rpc.mtls.AbstractMtlsTransportChannelTest; import com.google.api.gax.rpc.mtls.CertificateBasedAccess; +import com.google.api.gax.rpc.testing.FakeMtlsProvider; import com.google.auth.ApiKeyCredentials; import com.google.auth.Credentials; import com.google.auth.http.AuthHttpConstants; @@ -664,7 +665,9 @@ private void createAndCloseTransportChannel(InstantiatingGrpcChannelProvider pro createAndCloseTransportChannel(provider); assertThat(logHandler.getAllMessages()) .contains( - "DirectPath is misconfigured. The DirectPath XDS option was set, but the attemptDirectPath option was not. Please set both the attemptDirectPath and attemptDirectPathXds options."); + "DirectPath is misconfigured. The DirectPath XDS option was set, but the" + + " attemptDirectPath option was not. Please set both the attemptDirectPath and" + + " attemptDirectPathXds options."); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } @@ -682,8 +685,10 @@ void testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSe createAndCloseTransportChannel(provider); assertThat(logHandler.getAllMessages()) .contains( - "Env var GOOGLE_CLOUD_ENABLE_DIRECT_PATH_XDS was found and set to TRUE, but DirectPath was not enabled for this client. If this is intended for " - + "this client, please note that this is a misconfiguration and set the attemptDirectPath option as well."); + "Env var GOOGLE_CLOUD_ENABLE_DIRECT_PATH_XDS was found and set to TRUE, but DirectPath" + + " was not enabled for this client. If this is intended for this client, please" + + " note that this is a misconfiguration and set the attemptDirectPath option as" + + " well."); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } @@ -711,6 +716,7 @@ void testLogDirectPathMisconfigWrongCredential() throws Exception { InstantiatingGrpcChannelProvider.newBuilder() .setAttemptDirectPathXds() .setAttemptDirectPath(true) + .setEnvProvider(name -> null) .setHeaderProvider( mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) .setExecutor(mock(Executor.class)) @@ -877,12 +883,14 @@ public void canUseDirectPath_directPathEnvVarDisabled() throws IOException { @Test public void canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue() { System.setProperty("os.name", "Linux"); + EnvironmentProvider envProvider = name -> null; InstantiatingGrpcChannelProvider.Builder builder = InstantiatingGrpcChannelProvider.newBuilder() .setCertificateBasedAccess(certificateBasedAccess) .setAttemptDirectPath(true) .setCredentials(computeEngineCredentials) - .setEndpoint(DEFAULT_ENDPOINT); + .setEndpoint(DEFAULT_ENDPOINT) + .setEnvProvider(envProvider); InstantiatingGrpcChannelProvider provider = new InstantiatingGrpcChannelProvider(builder, GCE_PRODUCTION_NAME_AFTER_2016); Truth.assertThat(provider.canUseDirectPath()).isTrue(); @@ -891,12 +899,14 @@ public void canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue() { @Test public void canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsFalse() { System.setProperty("os.name", "Linux"); + EnvironmentProvider envProvider = name -> null; InstantiatingGrpcChannelProvider.Builder builder = InstantiatingGrpcChannelProvider.newBuilder() .setCertificateBasedAccess(certificateBasedAccess) .setAttemptDirectPath(false) .setCredentials(computeEngineCredentials) - .setEndpoint(DEFAULT_ENDPOINT); + .setEndpoint(DEFAULT_ENDPOINT) + .setEnvProvider(envProvider); InstantiatingGrpcChannelProvider provider = new InstantiatingGrpcChannelProvider(builder, GCE_PRODUCTION_NAME_AFTER_2016); Truth.assertThat(provider.canUseDirectPath()).isFalse(); @@ -1201,7 +1211,8 @@ void createS2ASecuredChannelCredentials_bothS2AAddressesNull_returnsNull() { assertThat(provider.createS2ASecuredChannelCredentials()).isNotNull(); assertThat(logHandler.getAllMessages()) .contains( - "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not return a mtls address to reach S2A."); + "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not return" + + " a mtls address to reach S2A."); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } @@ -1247,7 +1258,8 @@ void createS2ASecuredChannelCredentials_returnsPlaintextToS2AS2AChannelCredentia assertThat(provider.createS2ASecuredChannelCredentials()).isNotNull(); assertThat(logHandler.getAllMessages()) .contains( - "Cannot establish an mTLS connection to S2A because MTLS to MDS credentials do not exist on filesystem, falling back to plaintext connection to S2A"); + "Cannot establish an mTLS connection to S2A because MTLS to MDS credentials do not" + + " exist on filesystem, falling back to plaintext connection to S2A"); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } @@ -1342,6 +1354,96 @@ void testSettingBackgroundExecutor() { assertThat(provider.getBackgroundExecutor()).isEqualTo(mockExecutor); } + @Test + void createChannel_whenDirectPathEnabled_ignoresWorkloadCertPath() throws Exception { + System.setProperty("os.name", "Linux"); + EnvironmentProvider envProvider = + mock(EnvironmentProvider.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when( + envProvider.getenv( + InstantiatingGrpcChannelProvider.DIRECT_PATH_ENV_DISABLE_DIRECT_PATH)) + .thenReturn("false"); + CertificateBasedAccess mtlsCertificateBasedAccess = + mock(CertificateBasedAccess.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(mtlsCertificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(mtlsCertificateBasedAccess.getWorkloadCertPath()) + .thenReturn("/path/to/workload/cert.pem"); + MtlsProvider mtlsProvider = + new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingGrpcChannelProvider.Builder builder = + InstantiatingGrpcChannelProvider.newBuilder() + .setCertificateBasedAccess(mtlsCertificateBasedAccess) + .setMtlsProvider(mtlsProvider) + .setAttemptDirectPath(true) + .setCredentials(computeEngineCredentials) + .setEndpoint(DEFAULT_ENDPOINT) + .setEnvProvider(envProvider) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())); + InstantiatingGrpcChannelProvider provider = + new InstantiatingGrpcChannelProvider(builder, GCE_PRODUCTION_NAME_AFTER_2016); + Truth.assertThat(provider.canUseDirectPath()).isTrue(); + + TransportChannel transportChannel = provider.getTransportChannel(); + try { + ChannelPool pool = (ChannelPool) ((GrpcTransportChannel) transportChannel).getChannel(); + assertThat(pool.getWorkloadCertPath()).isNull(); + } finally { + transportChannel.shutdownNow(); + } + } + + @Test + void createChannel_whenMtlsActive_passesWorkloadCertPathToChannelPool() throws Exception { + CertificateBasedAccess mtlsCertificateBasedAccess = + mock(CertificateBasedAccess.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(mtlsCertificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(mtlsCertificateBasedAccess.getWorkloadCertPath()) + .thenReturn("/path/to/workload/cert.pem"); + MtlsProvider mtlsProvider = + new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingGrpcChannelProvider provider = + InstantiatingGrpcChannelProvider.newBuilder() + .setCertificateBasedAccess(mtlsCertificateBasedAccess) + .setMtlsProvider(mtlsProvider) + .setAttemptDirectPath(false) + .setEndpoint(DEFAULT_ENDPOINT) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) + .build(); + + TransportChannel transportChannel = provider.getTransportChannel(); + try { + ChannelPool pool = (ChannelPool) ((GrpcTransportChannel) transportChannel).getChannel(); + assertThat(pool.getWorkloadCertPath()).isEqualTo("/path/to/workload/cert.pem"); + } finally { + transportChannel.shutdownNow(); + } + } + + @Test + void createChannelBuilder_whenMtlsActiveAndKeyStoreNull_throwsIOException() { + CertificateBasedAccess mtlsCertificateBasedAccess = + mock(CertificateBasedAccess.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(mtlsCertificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + MtlsProvider mtlsProviderWithNullKeyStore = new FakeMtlsProvider(null, "", false); + + InstantiatingGrpcChannelProvider provider = + InstantiatingGrpcChannelProvider.newBuilder() + .setCertificateBasedAccess(mtlsCertificateBasedAccess) + .setMtlsProvider(mtlsProviderWithNullKeyStore) + .setAttemptDirectPath(false) + .setEndpoint(DEFAULT_ENDPOINT) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) + .build(); + + IOException thrown = assertThrows(IOException.class, provider::createChannelBuilder); + assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS channel credentials"); + } + private static class FakeLogHandler extends Handler { List records = new ArrayList<>(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java index 2679b51860df..81e970af4b06 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java @@ -82,6 +82,7 @@ public final class HttpJsonCallContext implements ApiCallContext { private final @Nullable RetrySettings retrySettings; private final @Nullable ImmutableSet retryableCodes; private final EndpointContext endpointContext; + @Nullable private final TransportChannel transportChannel; /** Returns an empty instance. */ public static HttpJsonCallContext createDefault() { @@ -96,6 +97,7 @@ public static HttpJsonCallContext createDefault() { null, null, null, + null, null); } @@ -111,6 +113,7 @@ public static HttpJsonCallContext of(HttpJsonChannel channel, HttpJsonCallOption null, null, null, + null, null); } @@ -125,7 +128,8 @@ private HttpJsonCallContext( @Nullable ApiTracer tracer, @Nullable RetrySettings defaultRetrySettings, @Nullable Set defaultRetryableCodes, - @Nullable EndpointContext endpointContext) { + @Nullable EndpointContext endpointContext, + @Nullable TransportChannel transportChannel) { this.channel = channel; this.callOptions = callOptions; this.timeout = timeout; @@ -141,6 +145,7 @@ private HttpJsonCallContext( // a valid EndpointContext with user configurations after the client has been initialized. this.endpointContext = endpointContext == null ? EndpointContext.getDefaultInstance() : endpointContext; + this.transportChannel = transportChannel; } /** @@ -220,6 +225,12 @@ public HttpJsonCallContext merge(ApiCallContext inputCallContext) { newRetryableCodes = this.retryableCodes; } + TransportChannel newTransportChannel = httpJsonCallContext.transportChannel; + if (newTransportChannel == null + && (httpJsonCallContext.channel == null || httpJsonCallContext.channel.equals(channel))) { + newTransportChannel = this.transportChannel; + } + // The EndpointContext is not updated as there should be no reason for a user // to update this. return new HttpJsonCallContext( @@ -233,7 +244,8 @@ public HttpJsonCallContext merge(ApiCallContext inputCallContext) { newTracer, newRetrySettings, newRetryableCodes, - endpointContext); + endpointContext, + newTransportChannel); } @Override @@ -251,7 +263,24 @@ public HttpJsonCallContext withTransportChannel(TransportChannel inputChannel) { "Expected HttpJsonTransportChannel, got " + inputChannel.getClass().getName()); } HttpJsonTransportChannel transportChannel = (HttpJsonTransportChannel) inputChannel; - return withChannel(transportChannel.getChannel()); + return new HttpJsonCallContext( + transportChannel.getChannel(), + this.callOptions, + this.timeout, + this.streamWaitTimeout, + this.streamIdleTimeout, + this.extraHeaders, + this.options, + this.tracer, + this.retrySettings, + this.retryableCodes, + this.endpointContext, + transportChannel); + } + + @Override + public TransportChannel getTransportChannel() { + return transportChannel; } /** This method is obsolete. Use {@link #withTimeoutDuration(java.time.Duration)} instead. */ @@ -275,7 +304,8 @@ public HttpJsonCallContext withEndpointContext(EndpointContext endpointContext) this.tracer, this.retrySettings, this.retryableCodes, - endpointContext); + endpointContext, + this.transportChannel); } @Override @@ -301,7 +331,8 @@ public HttpJsonCallContext withTimeoutDuration(java.time.Duration timeout) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** This method is obsolete. Use {@link #getTimeoutDuration()} instead. */ @@ -346,7 +377,8 @@ public HttpJsonCallContext withStreamWaitTimeoutDuration( this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** This method is obsolete. Use {@link #getStreamWaitTimeoutDuration()} instead. */ @@ -396,7 +428,8 @@ public HttpJsonCallContext withStreamIdleTimeoutDuration( this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** This method is obsolete. Use {@link #getStreamIdleTimeoutDuration()} instead. */ @@ -433,7 +466,8 @@ public ApiCallContext withExtraHeaders(Map> extraHeaders) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @BetaApi("The surface for extra headers is not stable yet and may change in the future.") @@ -457,7 +491,8 @@ public ApiCallContext withOption(Key key, T value) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** {@inheritDoc} */ @@ -527,7 +562,8 @@ public HttpJsonCallContext withRetrySettings(RetrySettings retrySettings) { this.tracer, retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @Override @@ -548,7 +584,8 @@ public HttpJsonCallContext withRetryableCodes(Set retryableCode this.tracer, this.retrySettings, retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } public HttpJsonCallContext withChannel(@Nullable HttpJsonChannel newChannel) { @@ -563,7 +600,8 @@ public HttpJsonCallContext withChannel(@Nullable HttpJsonChannel newChannel) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + (newChannel != null && newChannel.equals(this.channel)) ? this.transportChannel : null); } public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) { @@ -578,7 +616,8 @@ public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @Deprecated @@ -614,7 +653,8 @@ public HttpJsonCallContext withTracer(@Nonnull ApiTracer newTracer) { newTracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @Override @@ -634,7 +674,8 @@ public boolean equals(@Nullable Object o) { && Objects.equals(this.tracer, that.tracer) && Objects.equals(this.retrySettings, that.retrySettings) && Objects.equals(this.retryableCodes, that.retryableCodes) - && Objects.equals(this.endpointContext, that.endpointContext); + && Objects.equals(this.endpointContext, that.endpointContext) + && Objects.equals(this.transportChannel, that.transportChannel); } @Override @@ -648,6 +689,7 @@ public int hashCode() { tracer, retrySettings, retryableCodes, - endpointContext); + endpointContext, + transportChannel); } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java index 813622b6a97e..fde0673f650d 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java @@ -64,6 +64,21 @@ public HttpJsonChannel getChannel() { return getManagedChannel(); } + @Override + public void refresh() { + getManagedChannel().refresh(); + } + + @Override + public boolean shouldRefresh() { + return getManagedChannel().shouldRefresh(); + } + + @Override + public long getGeneration() { + return getManagedChannel().getGeneration(); + } + @Override public void shutdown() { getManagedChannel().shutdown(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 92ce4efe36aa..4ba35b01a8bd 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -233,33 +233,90 @@ private NetHttpTransport.Builder configureMtls(NetHttpTransport.Builder builder) return builder; } - private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecurityException { - HttpTransport httpTransportToUse = httpTransport; - if (httpTransportToUse == null) { - httpTransportToUse = createHttpTransport(); + private HttpTransport createChannelHttpTransport() throws IOException, GeneralSecurityException { + HttpTransport transport = createHttpTransport(); + if (mtlsProvider != null + && certificateBasedAccess.useMtlsClientCertificate() + && !((NetHttpTransport) transport).isMtls()) { + // mTLS is enabled but the provider returned no client certificate. Fail instead of silently + // using a transport without the certificate, matching InstantiatingGrpcChannelProvider. + // During certificate rotation, this makes RefreshingHttpJsonChannel keep the current + // authenticated transport instead of swapping in one without a client certificate. + throw new IOException("Failed to initialize mTLS HttpTransport"); } + return transport; + } + + private ManagedHttpJsonChannel createSingleManagedChannel() + throws IOException, GeneralSecurityException { + return buildManagedChannel( + httpTransport != null ? httpTransport : createChannelHttpTransport()); + } - // Pass the executor to the ManagedChannel. If no executor was provided (or null), - // the channel will use a default executor for the calls. - ManagedHttpJsonChannel channel = - ManagedHttpJsonChannel.newBuilder() - .setEndpoint(endpoint) - .setExecutor(executor) - .setHttpTransport(httpTransportToUse) - .build(); - - HttpJsonClientInterceptor headerInterceptor = - new HttpJsonHeaderInterceptor(headerProvider.getHeaders()); - - channel = new ManagedHttpJsonInterceptorChannel(channel, new HttpJsonLoggingInterceptor()); - channel = new ManagedHttpJsonInterceptorChannel(channel, headerInterceptor); - if (interceptorProvider != null && interceptorProvider.getInterceptors() != null) { - for (HttpJsonClientInterceptor interceptor : interceptorProvider.getInterceptors()) { - channel = new ManagedHttpJsonInterceptorChannel(channel, interceptor); + private ManagedHttpJsonChannel buildManagedChannel(HttpTransport httpTransportToUse) { + return ManagedHttpJsonChannel.newBuilder() + .setEndpoint(endpoint) + .setExecutor(executor) + .setHttpTransport(httpTransportToUse) + .setManageHttpTransport(httpTransport == null) + .build(); + } + + private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecurityException { + boolean isMtlsActive = + httpTransport == null + && mtlsProvider != null + && certificateBasedAccess.useMtlsClientCertificate(); + String workloadCertPath = isMtlsActive ? certificateBasedAccess.getWorkloadCertPath() : null; + + ManagedHttpJsonChannel baseChannel; + if (workloadCertPath != null) { + java.util.function.Supplier transportFactory = + () -> { + try { + return createChannelHttpTransport(); + } catch (IOException | GeneralSecurityException e) { + throw new java.lang.RuntimeException("Failed to create mTLS HttpTransport", e); + } + }; + // RefreshingHttpJsonChannel records the baseline certificate fingerprint before creating the + // initial transport, so a rotation during startup is detected on the next auth failure. + try { + baseChannel = + new RefreshingHttpJsonChannel( + transportFactory, this::buildManagedChannel, workloadCertPath); + } catch (RuntimeException e) { + if (e.getCause() instanceof IOException) { + throw (IOException) e.getCause(); + } + if (e.getCause() instanceof GeneralSecurityException) { + throw (GeneralSecurityException) e.getCause(); + } + throw e; } + } else { + baseChannel = createSingleManagedChannel(); } - return HttpJsonTransportChannel.newBuilder().setManagedChannel(channel).build(); + try { + ManagedHttpJsonChannel channel = baseChannel; + + HttpJsonClientInterceptor headerInterceptor = + new HttpJsonHeaderInterceptor(headerProvider.getHeaders()); + + channel = new ManagedHttpJsonInterceptorChannel(channel, new HttpJsonLoggingInterceptor()); + channel = new ManagedHttpJsonInterceptorChannel(channel, headerInterceptor); + if (interceptorProvider != null && interceptorProvider.getInterceptors() != null) { + for (HttpJsonClientInterceptor interceptor : interceptorProvider.getInterceptors()) { + channel = new ManagedHttpJsonInterceptorChannel(channel, interceptor); + } + } + + return HttpJsonTransportChannel.newBuilder().setManagedChannel(channel).build(); + } catch (Throwable t) { + baseChannel.shutdownNow(); + throw t; + } } /** The endpoint to be used for the channel. */ diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java index 87767bee5c7f..bf79e17d7d89 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java @@ -51,12 +51,35 @@ public class ManagedHttpJsonChannel implements HttpJsonChannel, BackgroundResour private final Executor executor; private final boolean usingDefaultExecutor; private final String endpoint; - private final HttpTransport httpTransport; + private volatile HttpTransport httpTransport; + private final boolean usingDefaultTransport; private final ScheduledExecutorService deadlineScheduledExecutorService; private boolean isTransportShutdown; protected ManagedHttpJsonChannel() { - this(null, true, null, null); + this(null, true, null, null, true); + } + + /** + * Constructor for subclasses that delegate all calls to a wrapped channel. The argument is + * unused; it only distinguishes this overload from {@link #ManagedHttpJsonChannel()}, which would + * otherwise allocate a transport and executor that the wrapper never uses or shuts down. + */ + ManagedHttpJsonChannel(boolean isDelegatingWrapper) { + this.executor = null; + this.usingDefaultExecutor = false; + this.endpoint = null; + this.httpTransport = null; + this.usingDefaultTransport = false; + this.deadlineScheduledExecutorService = null; + } + + /** + * Returns a monotonic generation counter tracking the number of successful refreshes or channel + * rotations performed by this channel. Always {@code 0} for channels that do not refresh. + */ + public long getGeneration() { + return 0; } String getEndpoint() { @@ -68,11 +91,20 @@ HttpTransport getHttpTransport() { return httpTransport; } + /** + * Replaces the transport used by calls created after this method returns. Calls that were already + * created keep using the transport they were created with. + */ + void setHttpTransport(HttpTransport httpTransport) { + this.httpTransport = httpTransport; + } + private ManagedHttpJsonChannel( @Nullable Executor executor, boolean usingDefaultExecutor, @Nullable String endpoint, - @Nullable HttpTransport httpTransport) { + @Nullable HttpTransport httpTransport, + boolean usingDefaultTransport) { this.executor = executor; this.usingDefaultExecutor = usingDefaultExecutor; this.endpoint = endpoint; @@ -82,6 +114,7 @@ private ManagedHttpJsonChannel( new NetHttpTransport.Builder()) .build() : httpTransport; + this.usingDefaultTransport = usingDefaultTransport || httpTransport == null; this.deadlineScheduledExecutorService = Executors.newSingleThreadScheduledExecutor(); } @@ -98,6 +131,22 @@ public HttpJsonClientCall newCall( deadlineScheduledExecutorService); } + /** + * Refreshes or recreates the underlying transport of this channel if a certificate rotation has + * been detected. By default, this is a no-op. + */ + public void refresh() { + // No-op: this channel has no certificate to rotate. Overridden by RefreshingHttpJsonChannel. + } + + /** + * Returns true if a certificate rotation has been detected on disk and this channel should be + * refreshed, or false otherwise. Always {@code false} for channels that do not refresh. + */ + public boolean shouldRefresh() { + return false; + } + @VisibleForTesting Executor getExecutor() { return executor; @@ -116,7 +165,9 @@ public synchronized void shutdown() { ((ExecutorService) executor).shutdown(); } deadlineScheduledExecutorService.shutdown(); - httpTransport.shutdown(); + if (usingDefaultTransport) { + httpTransport.shutdown(); + } isTransportShutdown = true; } catch (IOException e) { // TODO: Log this scenario once we implemented the Cloud SDK logging. @@ -158,7 +209,9 @@ public void shutdownNow() { ((ExecutorService) executor).shutdownNow(); } deadlineScheduledExecutorService.shutdownNow(); - httpTransport.shutdown(); + if (usingDefaultTransport) { + httpTransport.shutdown(); + } isTransportShutdown = true; } catch (IOException e) { // TODO: Log this scenario once we implemented the Cloud SDK logging. @@ -205,9 +258,11 @@ public static class Builder { private String endpoint; private HttpTransport httpTransport; private boolean usingDefaultExecutor; + private boolean usingDefaultTransport; private Builder() { this.usingDefaultExecutor = false; + this.usingDefaultTransport = false; } public Builder setExecutor(Executor executor) { @@ -225,6 +280,11 @@ public Builder setHttpTransport(HttpTransport httpTransport) { return this; } + Builder setManageHttpTransport(boolean manageHttpTransport) { + this.usingDefaultTransport = manageHttpTransport; + return this; + } + public ManagedHttpJsonChannel build() { Preconditions.checkNotNull(endpoint); @@ -237,14 +297,8 @@ public ManagedHttpJsonChannel build() { usingDefaultExecutor = true; } - if (httpTransport == null) { - httpTransport = - HttpJsonConscryptUtils.configureConscryptSecurityProvider( - new NetHttpTransport.Builder()) - .build(); - } - - return new ManagedHttpJsonChannel(executor, usingDefaultExecutor, endpoint, httpTransport); + return new ManagedHttpJsonChannel( + executor, usingDefaultExecutor, endpoint, httpTransport, usingDefaultTransport); } } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java index eaaa8c3a7c56..9adbfdd5c329 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java @@ -29,7 +29,9 @@ */ package com.google.api.gax.httpjson; +import com.google.api.client.http.HttpTransport; import com.google.common.annotations.VisibleForTesting; +import java.util.concurrent.Executor; import java.util.concurrent.TimeUnit; import org.jspecify.annotations.NullMarked; @@ -41,22 +43,57 @@ class ManagedHttpJsonInterceptorChannel extends ManagedHttpJsonChannel { ManagedHttpJsonInterceptorChannel( ManagedHttpJsonChannel channel, HttpJsonClientInterceptor interceptor) { - super(); + super(true); this.channel = channel; this.interceptor = interceptor; } + /** {@inheritDoc} */ + @Override + public long getGeneration() { + return channel.getGeneration(); + } + @VisibleForTesting ManagedHttpJsonChannel getChannel() { return channel; } + @Override + String getEndpoint() { + return channel.getEndpoint(); + } + + @Override + @VisibleForTesting + HttpTransport getHttpTransport() { + return channel.getHttpTransport(); + } + + @Override + @VisibleForTesting + Executor getExecutor() { + return channel.getExecutor(); + } + @Override public HttpJsonClientCall newCall( ApiMethodDescriptor methodDescriptor, HttpJsonCallOptions callOptions) { return interceptor.interceptCall(methodDescriptor, callOptions, channel); } + /** {@inheritDoc} */ + @Override + public void refresh() { + channel.refresh(); + } + + /** {@inheritDoc} */ + @Override + public boolean shouldRefresh() { + return channel.shouldRefresh(); + } + @Override public synchronized void shutdown() { channel.shutdown(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java new file mode 100644 index 000000000000..6d20a69139f8 --- /dev/null +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -0,0 +1,213 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.httpjson; + +import com.google.api.client.http.HttpTransport; +import com.google.api.core.InternalApi; +import com.google.api.gax.rpc.mtls.CertificateRotationTracker; +import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils; +import com.google.common.annotations.VisibleForTesting; +import java.util.concurrent.Executor; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicLong; +import java.util.function.Function; +import java.util.function.Supplier; +import java.util.logging.Level; +import java.util.logging.Logger; + +/** + * An implementation of {@link ManagedHttpJsonChannel} that supports dynamic mTLS certificate + * rotation. When the workload certificate on disk changes, {@link #refresh()} replaces the {@link + * HttpTransport} of the underlying channel. Calls that were already created keep using the + * transport they were created with, so in-flight requests are not interrupted. + */ +@InternalApi +public class RefreshingHttpJsonChannel extends ManagedHttpJsonChannel { + + private static final Logger LOG = Logger.getLogger(RefreshingHttpJsonChannel.class.getName()); + + private final CertificateRotationTracker rotationTracker; + private final Supplier transportFactory; + private final String workloadCertPath; + private final ManagedHttpJsonChannel delegate; + private final Object refreshLock = new Object(); + private final AtomicLong generation = new AtomicLong(0); + + /** + * @param transportFactory creates a transport configured with the certificate currently on disk + * @param channelFactory creates the underlying channel from the initial transport + * @param workloadCertPath path of the workload certificate to monitor for rotation + */ + public RefreshingHttpJsonChannel( + Supplier transportFactory, + Function channelFactory, + String workloadCertPath) { + super(true); + this.transportFactory = transportFactory; + this.workloadCertPath = workloadCertPath; + // Record the baseline fingerprint before the initial transport loads the certificate from disk, + // so a rotation between the two steps is detected as a fingerprint change. + this.rotationTracker = + new CertificateRotationTracker(this::getWorkloadCertPath, this::getCertificateFingerprint); + this.delegate = channelFactory.apply(transportFactory.get()); + } + + @VisibleForTesting + String getWorkloadCertPath() { + return workloadCertPath; + } + + @VisibleForTesting + String getCertificateFingerprint(String certPath) { + return WorkloadCertificateUtils.getCertificateFingerprint(certPath); + } + + /** {@inheritDoc} */ + @Override + public boolean shouldRefresh() { + return rotationTracker.shouldRefresh(); + } + + /** {@inheritDoc} */ + @Override + public void refresh() { + // A generation change while waiting for the lock means a concurrent refresh already swapped in + // a transport with a new certificate, so there is no need to read the certificate again. + long generationBeforeLock = generation.get(); + synchronized (refreshLock) { + if (isShutdown()) { + return; + } + if (generation.get() != generationBeforeLock) { + LOG.fine( + "HTTP/JSON channel was already refreshed by a concurrent thread, skipping duplicate" + + " refresh"); + return; + } + String currentDiskFingerprint = rotationTracker.readDiskFingerprint(); + if (currentDiskFingerprint.isEmpty()) { + return; + } + + // Double-check inside refreshLock + if (rotationTracker.isAlreadyActive(currentDiskFingerprint)) { + LOG.fine( + "HTTP/JSON channel was already refreshed by a concurrent thread, skipping duplicate" + + " refresh"); + return; + } + + LOG.info("mTLS certificate rotation detected. Refreshing HTTP/JSON transport."); + + HttpTransport newTransport; + try { + newTransport = transportFactory.get(); + } catch (Exception e) { + LOG.log(Level.WARNING, "Failed to refresh HTTP/JSON transport, keeping old transport", e); + return; + } + // The previous transport is not shut down: calls created before the swap may still be using + // it, and it needs no explicit shutdown because its idle keep-alive connections expire. + delegate.setHttpTransport(newTransport); + // Order matters: swap the transport, then bump generation, then mark the tracker refreshed, + // so any failing RPC that observes the new fingerprint also observes the new generation. + generation.incrementAndGet(); + rotationTracker.markRefreshed(currentDiskFingerprint); + } + } + + /** {@inheritDoc} */ + @Override + public long getGeneration() { + return generation.get(); + } + + @Override + public HttpJsonClientCall newCall( + ApiMethodDescriptor methodDescriptor, HttpJsonCallOptions callOptions) { + return delegate.newCall(methodDescriptor, callOptions); + } + + @Override + Executor getExecutor() { + return delegate.getExecutor(); + } + + @Override + String getEndpoint() { + return delegate.getEndpoint(); + } + + @Override + @VisibleForTesting + HttpTransport getHttpTransport() { + return delegate.getHttpTransport(); + } + + @VisibleForTesting + void invalidateDiskFingerprintCache() { + rotationTracker.invalidateCache(); + } + + @Override + public void shutdown() { + // Serialized with refresh() so that a transport swap cannot race with shutdown. + synchronized (refreshLock) { + delegate.shutdown(); + } + } + + @Override + public boolean isShutdown() { + return delegate.isShutdown(); + } + + @Override + public boolean isTerminated() { + return delegate.isTerminated(); + } + + @Override + public void shutdownNow() { + synchronized (refreshLock) { + delegate.shutdownNow(); + } + } + + @Override + public boolean awaitTermination(long duration, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(duration, unit); + } + + @Override + public void close() { + shutdown(); + } +} diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java index 08044522e729..39c82205c294 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java @@ -207,6 +207,20 @@ void testWithLongerTimeout() { .isEqualTo(java.time.Duration.ofSeconds(5)); } + @Test + void testWithNullOrZeroTimeoutClearsExistingTimeout() { + HttpJsonCallContext ctxWithTimeout = + HttpJsonCallContext.createDefault().withTimeoutDuration(java.time.Duration.ofSeconds(5)); + + // Sanity check + Truth.assertThat(ctxWithTimeout.getTimeoutDuration()) + .isEqualTo(java.time.Duration.ofSeconds(5)); + + java.time.Duration nullTimeout = null; + assertNull(ctxWithTimeout.withTimeoutDuration(nullTimeout).getTimeoutDuration()); + assertNull(ctxWithTimeout.withTimeoutDuration(java.time.Duration.ZERO).getTimeoutDuration()); + } + @Test void testMergeWithNullTimeout() { java.time.Duration timeout = java.time.Duration.ofSeconds(10); @@ -334,4 +348,55 @@ void testMergeOptions() { assertEquals(testContext2, mergedContext.getOption(contextKey2)); assertEquals(testContext3, mergedContext.getOption(contextKey3)); } + + @Test + void testWithChannelClearsStaleTransportChannel() { + ManagedHttpJsonChannel channel1 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + ManagedHttpJsonChannel channel2 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + + HttpJsonTransportChannel transportChannel1 = + HttpJsonTransportChannel.newBuilder().setManagedChannel(channel1).build(); + + HttpJsonCallContext context = + HttpJsonCallContext.createDefault().withTransportChannel(transportChannel1); + Truth.assertThat(context.getTransportChannel()).isSameInstanceAs(transportChannel1); + + // Retains transportChannel when setting same channel + Truth.assertThat(context.withChannel(channel1).getTransportChannel()) + .isSameInstanceAs(transportChannel1); + + // Clears transportChannel to null when setting null or a different channel + HttpJsonCallContext nullChannelContext = context.withChannel(null); + Truth.assertThat(nullChannelContext.getChannel()).isNull(); + Truth.assertThat(nullChannelContext.getTransportChannel()).isNull(); + Truth.assertThat(context.withChannel(channel2).getTransportChannel()).isNull(); + + // Merging a context with a cleared channel into the original context keeps the original + // channel and transportChannel + HttpJsonCallContext mergedWithNullChannel = context.merge(nullChannelContext); + Truth.assertThat(mergedWithNullChannel.getChannel()).isSameInstanceAs(channel1); + Truth.assertThat(mergedWithNullChannel.getTransportChannel()) + .isSameInstanceAs(transportChannel1); + } + + @Test + void testMergeClearsStaleTransportChannel() { + ManagedHttpJsonChannel channel1 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + ManagedHttpJsonChannel channel2 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + + HttpJsonTransportChannel transportChannel1 = + HttpJsonTransportChannel.newBuilder().setManagedChannel(channel1).build(); + + HttpJsonCallContext context1 = + HttpJsonCallContext.createDefault().withTransportChannel(transportChannel1); + HttpJsonCallContext context2 = HttpJsonCallContext.createDefault().withChannel(channel2); + + HttpJsonCallContext merged = context1.merge(context2); + Truth.assertThat(merged.getChannel()).isSameInstanceAs(channel2); + Truth.assertThat(merged.getTransportChannel()).isNull(); + } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index 8c95c1d2e1c4..ee9c476adf18 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -31,13 +31,16 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.mockito.Mockito.mock; import com.google.api.client.http.javanet.NetHttpTransport; +import com.google.api.gax.rpc.HeaderProvider; import com.google.api.gax.rpc.TransportChannelProvider; import com.google.api.gax.rpc.mtls.AbstractMtlsTransportChannelTest; import com.google.api.gax.rpc.mtls.CertificateBasedAccess; import com.google.auth.mtls.MtlsProvider; import java.io.IOException; +import java.nio.charset.StandardCharsets; import java.security.GeneralSecurityException; import java.util.Collections; import java.util.Map; @@ -46,6 +49,7 @@ import java.util.concurrent.ScheduledThreadPoolExecutor; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.mockito.Mockito; class InstantiatingHttpJsonChannelProviderTest extends AbstractMtlsTransportChannelTest { @@ -55,9 +59,10 @@ class InstantiatingHttpJsonChannelProviderTest extends AbstractMtlsTransportChan @BeforeEach public void setup() throws IOException { - certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "never" : "false"); + certificateBasedAccess = org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(false); } @Test @@ -179,20 +184,361 @@ void managedChannelUsesCustomExecutor() throws IOException { instantiatingHttpJsonChannelProvider.getTransportChannel().shutdownNow(); } - @Override - protected Object getMtlsObjectFromTransportChannel( - MtlsProvider provider, CertificateBasedAccess certificateBasedAccess) + @Test + void managedChannelDoesNotShutdownCustomHttpTransport() throws IOException { + com.google.api.client.http.HttpTransport mockHttpTransport = + org.mockito.Mockito.mock(com.google.api.client.http.HttpTransport.class); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setHttpTransport(mockHttpTransport) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + + // Verify custom transport is injected (direct ManagedHttpJsonChannel when workloadCertPath is + // null) + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonInterceptorChannel managedHttpJsonChannel = + (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel(); + ManagedHttpJsonChannel channel = managedHttpJsonChannel.getChannel(); + + assertThat(channel.getHttpTransport()).isEqualTo(mockHttpTransport); + + // Perform a shutdown + provider.getTransportChannel().shutdownNow(); + + // Verify that shutdown() was NOT called on the custom HttpTransport + org.mockito.Mockito.verify(mockHttpTransport, org.mockito.Mockito.never()).shutdown(); + } + + @Test + void channelCreation_withWorkloadCertPath_wrapsWithRefreshingHttpJsonChannel() + throws IOException, GeneralSecurityException { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(mtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonInterceptorChannel managedHttpJsonChannel = + (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel(); + assertThat(managedHttpJsonChannel.getChannel()).isInstanceOf(RefreshingHttpJsonChannel.class); + + provider.getTransportChannel().shutdownNow(); + } + + @Test + void channelCreation_withCustomHttpTransport_ignoresWorkloadCertPathAndDoesNotWrap() + throws IOException, GeneralSecurityException { + // mTLS is otherwise fully configured, so only the custom transport prevents wrapping. Lenient + // because a custom transport is expected to skip these lookups. + Mockito.lenient().when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.lenient() + .when(certificateBasedAccess.getWorkloadCertPath()) + .thenReturn("fake/cert/path.json"); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + com.google.api.client.http.HttpTransport mockHttpTransport = + org.mockito.Mockito.mock(com.google.api.client.http.HttpTransport.class); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setHttpTransport(mockHttpTransport) + .setMtlsProvider(mtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonInterceptorChannel managedHttpJsonChannel = + (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel(); + assertThat(managedHttpJsonChannel.getChannel()) + .isNotInstanceOf(RefreshingHttpJsonChannel.class); + Mockito.verify(certificateBasedAccess, Mockito.never()).getWorkloadCertPath(); + + httpJsonTransportChannel.shutdownNow(); + } + + @Test + void getTransportChannel_withMtlsKeyStore_usesMtlsTransport() throws IOException { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(mtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonChannel channel = + ((ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel()).getChannel(); + + assertThat(channel).isNotInstanceOf(RefreshingHttpJsonChannel.class); + assertThat(((NetHttpTransport) channel.getHttpTransport()).isMtls()).isTrue(); + + httpJsonTransportChannel.shutdownNow(); + } + + @Test + void refresh_whenKeyStoreUnavailableDuringRotation_keepsCurrentTransport( + @org.junit.jupiter.api.io.TempDir java.nio.file.Path tempDir) throws IOException { + java.nio.file.Path certPath = tempDir.resolve("cert.pem"); + java.nio.file.Files.write(certPath, "certificate-v1".getBytes(StandardCharsets.UTF_8)); + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn(certPath.toString()); + java.security.KeyStore keyStore = + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(); + java.util.concurrent.atomic.AtomicReference currentKeyStore = + new java.util.concurrent.atomic.AtomicReference<>(keyStore); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.auth.mtls.MtlsProvider() { + @Override + public java.security.KeyStore getKeyStore() { + return currentKeyStore.get(); + } + + @Override + public boolean isAvailable() { + return true; + } + }; + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(mtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + RefreshingHttpJsonChannel channel = + (RefreshingHttpJsonChannel) + ((ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel()).getChannel(); + com.google.api.client.http.HttpTransport initialTransport = channel.getHttpTransport(); + assertThat(((NetHttpTransport) initialTransport).isMtls()).isTrue(); + + // The certificate rotates on disk, but the provider cannot return the new key store yet. + java.nio.file.Files.write(certPath, "certificate-v2".getBytes(StandardCharsets.UTF_8)); + currentKeyStore.set(null); + assertThat(channel.shouldRefresh()).isTrue(); + channel.refresh(); + + // The current authenticated transport is kept instead of one without a client certificate. + assertThat(channel.getHttpTransport()).isSameInstanceAs(initialTransport); + assertThat(channel.getGeneration()).isEqualTo(0); + assertThat(channel.shouldRefresh()).isTrue(); + + // Once the key store is available again, the next refresh swaps in a new mTLS transport. + currentKeyStore.set(keyStore); + channel.refresh(); + assertThat(channel.getHttpTransport()).isNotSameInstanceAs(initialTransport); + assertThat(((NetHttpTransport) channel.getHttpTransport()).isMtls()).isTrue(); + assertThat(channel.getGeneration()).isEqualTo(1); + + httpJsonTransportChannel.shutdownNow(); + } + + @Test + void getTransportChannel_whenMtlsKeyStoreThrowsIOException_throwsCheckedIOException() + throws Exception { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + MtlsProvider failingMtlsProvider = + Mockito.mock(MtlsProvider.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(failingMtlsProvider.getKeyStore()) + .thenThrow(new IOException("Simulated keystore read failure")); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(failingMtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + final InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + // Must throw checked IOException directly (not wrapped in RuntimeException) + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasMessageThat().contains("Simulated keystore read failure"); + } + + @Test + void getTransportChannel_whenMtlsActiveAndKeyStoreNull_throwsIOException() { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + com.google.auth.mtls.MtlsProvider providerWithNullKeyStore = + new com.google.api.gax.rpc.testing.FakeMtlsProvider(null, "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(providerWithNullKeyStore) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS HttpTransport"); + } + + @Test + void getTransportChannel_whenRotationEnabledAndKeyStoreNull_throwsIOException() { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + com.google.auth.mtls.MtlsProvider providerWithNullKeyStore = + new com.google.api.gax.rpc.testing.FakeMtlsProvider(null, "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(providerWithNullKeyStore) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS HttpTransport"); + } + + @Test + void getTransportChannel_whenKeyStoreUninitialized_wrapsGeneralSecurityException() + throws Exception { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + // A KeyStore that was never loaded makes transport creation fail with a KeyStoreException. + java.security.KeyStore uninitializedKeyStore = java.security.KeyStore.getInstance("PKCS12"); + com.google.auth.mtls.MtlsProvider providerWithBadKeyStore = + new com.google.api.gax.rpc.testing.FakeMtlsProvider(uninitializedKeyStore, "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(providerWithBadKeyStore) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + // The GeneralSecurityException must be unwrapped from the factory's RuntimeException, so that + // getTransportChannel() surfaces it as the direct cause of its checked IOException. + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasCauseThat().isInstanceOf(GeneralSecurityException.class); + } + + @Test + void getTransportChannel_whenMtlsKeyStoreThrowsRuntimeException_propagatesUnwrapped() + throws Exception { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + IllegalStateException failure = new IllegalStateException("Simulated keystore failure"); + MtlsProvider failingMtlsProvider = + Mockito.mock(MtlsProvider.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(failingMtlsProvider.getKeyStore()).thenThrow(failure); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(failingMtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + IllegalStateException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IllegalStateException.class, finalProvider::getTransportChannel); + assertThat(thrown).isSameInstanceAs(failure); + } + + @Test + void createHttpTransport_withMtlsKeyStore_returnsMtlsTransport() throws IOException, GeneralSecurityException { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + com.google.auth.mtls.MtlsProvider provider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + InstantiatingHttpJsonChannelProvider channelProvider = InstantiatingHttpJsonChannelProvider.newBuilder() - .setEndpoint("localhost:8080") + .setEndpoint(DEFAULT_ENDPOINT) .setMtlsProvider(provider) .setCertificateBasedAccess(certificateBasedAccess) - .setHeaderProvider(Collections::emptyMap) - .setExecutor(Runnable::run) .build(); - NetHttpTransport transport = (NetHttpTransport) channelProvider.createHttpTransport(); - return (transport != null && transport.isMtls()) ? transport : null; + + com.google.api.client.http.HttpTransport transport = channelProvider.createHttpTransport(); + assertThat(transport).isNotNull(); + assertThat(transport).isInstanceOf(com.google.api.client.http.javanet.NetHttpTransport.class); + assertThat(((com.google.api.client.http.javanet.NetHttpTransport) transport).isMtls()).isTrue(); + } + + @Test + void createHttpTransport_whenMtlsProviderNullOrNotUsingClientCert_returnsNonMtlsTransport() + throws IOException, GeneralSecurityException { + InstantiatingHttpJsonChannelProvider nullMtlsProviderChannelProvider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(null) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + NetHttpTransport nullMtlsProviderTransport = + (NetHttpTransport) nullMtlsProviderChannelProvider.createHttpTransport(); + assertThat(nullMtlsProviderTransport).isNotNull(); + assertThat(nullMtlsProviderTransport.isMtls()).isFalse(); + + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(false); + com.google.auth.mtls.MtlsProvider provider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + InstantiatingHttpJsonChannelProvider disabledMtlsChannelProvider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(provider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + NetHttpTransport disabledMtlsTransport = + (NetHttpTransport) disabledMtlsChannelProvider.createHttpTransport(); + assertThat(disabledMtlsTransport).isNotNull(); + assertThat(disabledMtlsTransport.isMtls()).isFalse(); } @Test @@ -214,4 +560,21 @@ void testConfigureConscryptSecurityProvider_returnsConfiguredBuilder() { HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); assertThat(result).isSameInstanceAs(builder); } + + @Override + protected Object getMtlsObjectFromTransportChannel( + MtlsProvider provider, CertificateBasedAccess certificateBasedAccess) + throws IOException, GeneralSecurityException { + InstantiatingHttpJsonChannelProvider channelProvider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint("localhost:8080") + .setMtlsProvider(provider) + .setCertificateBasedAccess(certificateBasedAccess) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) + .setExecutor(mock(Executor.class)) + .build(); + NetHttpTransport transport = (NetHttpTransport) channelProvider.createHttpTransport(); + return (transport != null && transport.isMtls()) ? transport : null; + } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java new file mode 100644 index 000000000000..1018005f2df5 --- /dev/null +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -0,0 +1,607 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.httpjson; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotSame; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import com.google.api.client.http.HttpTransport; +import com.google.api.client.testing.http.MockHttpTransport; +import com.google.api.gax.httpjson.testing.MockHttpService; +import com.google.protobuf.Field; +import java.util.ArrayDeque; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Queue; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.function.Function; +import java.util.function.Supplier; +import javax.annotation.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +class RefreshingHttpJsonChannelTest { + private static final ApiMethodDescriptor FAKE_METHOD_DESCRIPTOR = + ApiMethodDescriptor.newBuilder() + .setFullMethodName("google.cloud.v1.Fake/FakeMethod") + .setHttpMethod("POST") + .setRequestFormatter( + ProtoMessageRequestFormatter.newBuilder() + .setPath( + "/fake/v1/name/{name}", + request -> { + Map fields = new HashMap<>(); + ProtoRestSerializer serializer = ProtoRestSerializer.create(); + serializer.putPathParam(fields, "name", request.getName()); + return fields; + }) + .setQueryParamsExtractor(request -> new HashMap<>()) + .setRequestBodyExtractor( + request -> + ProtoRestSerializer.create() + .toBody("*", request.toBuilder().clearName().build(), false)) + .build()) + .setResponseParser( + ProtoMessageResponseParser.newBuilder() + .setDefaultInstance(Field.getDefaultInstance()) + .build()) + .build(); + + private static class FakeHttpJsonClientCall + extends HttpJsonClientCall { + @Override + public void start(Listener responseListener, HttpJsonMetadata requestHeaders) {} + + @Override + public void request(int numMessages) {} + + @Override + public void cancel(@Nullable String message, @Nullable Throwable cause) {} + + @Override + public void sendMessage(RequestT message) {} + + @Override + public void halfClose() {} + } + + private static class FakeManagedHttpJsonChannel extends ManagedHttpJsonChannel { + private volatile boolean isShutdown = false; + private volatile boolean isTerminated = false; + private HttpJsonClientCall nextCall = null; + + @Override + String getEndpoint() { + return "https://fake.endpoint:443"; + } + + @Override + public void shutdown() { + isShutdown = true; + } + + @Override + public void shutdownNow() { + isShutdown = true; + isTerminated = true; + } + + @Override + public boolean isShutdown() { + return isShutdown; + } + + @Override + public boolean isTerminated() { + return isTerminated; + } + + @Override + public boolean awaitTermination(long duration, TimeUnit unit) { + return isTerminated; + } + + @Override + @SuppressWarnings("unchecked") + public HttpJsonClientCall newCall( + ApiMethodDescriptor methodDescriptor, + HttpJsonCallOptions callOptions) { + if (nextCall != null) { + return (HttpJsonClientCall) nextCall; + } + return new FakeHttpJsonClientCall<>(); + } + } + + private AtomicInteger transportFactoryCount; + private HttpTransport lastCreatedTransport; + private FakeManagedHttpJsonChannel lastCreatedChannel; + private String testCertPath; + private String testFingerprint; + private boolean shouldThrowOnFactory; + private List createdChannels; + + private Supplier transportFactory = + () -> { + if (shouldThrowOnFactory) { + throw new RuntimeException("Simulated factory failure"); + } + transportFactoryCount.incrementAndGet(); + lastCreatedTransport = new MockHttpTransport(); + return lastCreatedTransport; + }; + + private Function channelFactory = + transport -> { + lastCreatedChannel = new FakeManagedHttpJsonChannel(); + lastCreatedChannel.setHttpTransport(transport); + return lastCreatedChannel; + }; + + @BeforeEach + void setUp() { + transportFactoryCount = new AtomicInteger(0); + testCertPath = "/fake/path"; + testFingerprint = "fingerprint1"; + shouldThrowOnFactory = false; + createdChannels = new ArrayList<>(); + } + + @AfterEach + void tearDown() { + for (RefreshingHttpJsonChannel channel : createdChannels) { + channel.shutdownNow(); + } + } + + private RefreshingHttpJsonChannel createTestChannel() { + RefreshingHttpJsonChannel ch = + new RefreshingHttpJsonChannel( + () -> transportFactory.get(), + transport -> channelFactory.apply(transport), + "fake/cert/path.json") { + @Override + String getWorkloadCertPath() { + return testCertPath; + } + + @Override + String getCertificateFingerprint(String certPath) { + return testFingerprint; + } + }; + createdChannels.add(ch); + return ch; + } + + private void rotateCertificate(RefreshingHttpJsonChannel channel) { + channel.invalidateDiskFingerprintCache(); + testFingerprint = "fingerprint2"; + } + + @Test + void testShouldRefreshFalseWhenCertPathNull() { + testCertPath = null; + RefreshingHttpJsonChannel channel = createTestChannel(); + assertFalse(channel.shouldRefresh()); + } + + @Test + void testShouldRefreshFalseWhenUnchanged() { + RefreshingHttpJsonChannel channel = createTestChannel(); + + channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache + assertFalse(channel.shouldRefresh()); + } + + @Test + void testShouldRefreshTrueWhenChanged() { + RefreshingHttpJsonChannel channel = createTestChannel(); + + rotateCertificate(channel); + + assertTrue(channel.shouldRefresh()); + } + + @Test + void shouldRefresh_doesNotCacheNegativeResultAndDetectsSubsequentRotationImmediately() { + RefreshingHttpJsonChannel channel = createTestChannel(); + + // First check returns false (unchanged fingerprint) + assertFalse(channel.shouldRefresh()); + + // Immediately change fingerprint WITHOUT invalidating cache + testFingerprint = "fingerprint2"; + + // Must immediately detect rotation because negative/unchanged checks are not cached for 1s + assertTrue(channel.shouldRefresh()); + + // Refresh updates activeCertFingerprint and clears cache + channel.refresh(); + assertFalse(channel.shouldRefresh()); + } + + @Test + void rotationDuringInitialTransportCreation_isDetectedAndRefreshed() { + transportFactory = + () -> { + if (transportFactoryCount.incrementAndGet() == 1) { + // Simulate the certificate rotating on disk while the initial transport loads it. + testFingerprint = "fingerprint2"; + } + lastCreatedTransport = new MockHttpTransport(); + return lastCreatedTransport; + }; + + RefreshingHttpJsonChannel channel = createTestChannel(); + + // The baseline was recorded before the initial transport was created, so the rotation is seen. + assertTrue(channel.shouldRefresh()); + + channel.refresh(); + assertEquals(2, transportFactoryCount.get()); + assertEquals(1, channel.getGeneration()); + assertFalse(channel.shouldRefresh()); + } + + @Test + void refresh_swapsTransportAndKeepsChannel() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel; + HttpTransport initialTransport = channel.getHttpTransport(); + assertSame(lastCreatedTransport, initialTransport); + + rotateCertificate(channel); + channel.refresh(); + + assertEquals(2, transportFactoryCount.get()); + assertNotSame(initialTransport, channel.getHttpTransport()); + assertSame(lastCreatedTransport, channel.getHttpTransport()); + // The underlying channel (and its executors) is reused rather than replaced or shut down. + assertSame(underlyingChannel, lastCreatedChannel); + assertFalse(underlyingChannel.isShutdown()); + assertEquals(1, channel.getGeneration()); + assertFalse(channel.shouldRefresh()); + } + + @Test + void refresh_whenCertificateFileEmpty_keepsTransportAndGeneration() { + RefreshingHttpJsonChannel channel = createTestChannel(); + HttpTransport initialTransport = channel.getHttpTransport(); + + // A rotation is detected, but when refresh() re-reads the file the rotator has truncated it + // and not yet written the new certificate. + rotateCertificate(channel); + assertTrue(channel.shouldRefresh()); + testFingerprint = ""; + channel.refresh(); + + // Nothing is swapped and the generation is unchanged, so the failed call is not retried. + assertEquals(1, transportFactoryCount.get()); + assertSame(initialTransport, channel.getHttpTransport()); + assertEquals(0, channel.getGeneration()); + // While the file stays empty, no refresh is requested. + channel.invalidateDiskFingerprintCache(); + assertFalse(channel.shouldRefresh()); + + // Once the new certificate is written, the next check refreshes normally. + testFingerprint = "fingerprint2"; + assertTrue(channel.shouldRefresh()); + channel.refresh(); + assertEquals(2, transportFactoryCount.get()); + assertNotSame(initialTransport, channel.getHttpTransport()); + assertEquals(1, channel.getGeneration()); + } + + @Test + void callCreatedBeforeRefresh_usesOriginalTransport() throws Exception { + MockHttpService originalService = + new MockHttpService(Collections.singletonList(FAKE_METHOD_DESCRIPTOR), "google.com:443"); + MockHttpService rotatedService = + new MockHttpService(Collections.singletonList(FAKE_METHOD_DESCRIPTOR), "google.com:443"); + Field message = Field.newBuilder().setName("bob").setNumber(1).build(); + originalService.addResponse(message); + rotatedService.addResponse(message); + Queue transports = + new ArrayDeque<>(Arrays.asList(originalService, rotatedService)); + transportFactory = transports::remove; + channelFactory = + transport -> + ManagedHttpJsonChannel.newBuilder() + .setEndpoint("google.com:443") + .setHttpTransport(transport) + .build(); + RefreshingHttpJsonChannel channel = createTestChannel(); + HttpJsonCallOptions callOptions = HttpJsonCallOptions.newBuilder().build(); + HttpJsonCallContext callContext = HttpJsonCallContext.createDefault(); + + HttpJsonClientCall callBeforeRefresh = + channel.newCall(FAKE_METHOD_DESCRIPTOR, callOptions); + + rotateCertificate(channel); + channel.refresh(); + assertEquals(1, channel.getGeneration()); + + assertEquals( + message, + HttpJsonClientCalls.futureUnaryCall(callBeforeRefresh, message, callContext) + .get(10, TimeUnit.SECONDS)); + assertEquals(1, originalService.getRequestPaths().size()); + assertEquals(0, rotatedService.getRequestPaths().size()); + + HttpJsonClientCall callAfterRefresh = + channel.newCall(FAKE_METHOD_DESCRIPTOR, callOptions); + assertEquals( + message, + HttpJsonClientCalls.futureUnaryCall(callAfterRefresh, message, callContext) + .get(10, TimeUnit.SECONDS)); + assertEquals(1, originalService.getRequestPaths().size()); + assertEquals(1, rotatedService.getRequestPaths().size()); + } + + @Test + void testRefreshDoesNotCreateTransportWhenShutdown() { + RefreshingHttpJsonChannel channel = createTestChannel(); + assertEquals(1, transportFactoryCount.get()); + + channel.shutdown(); + rotateCertificate(channel); + channel.refresh(); + + assertEquals(1, transportFactoryCount.get()); + assertEquals(0, channel.getGeneration()); + } + + @Test + void shutdown_waitsForInProgressRefresh() throws Exception { + CountDownLatch refreshStarted = new CountDownLatch(1); + CountDownLatch releaseRefresh = new CountDownLatch(1); + transportFactory = + () -> { + if (transportFactoryCount.incrementAndGet() > 1) { + refreshStarted.countDown(); + try { + releaseRefresh.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + return new MockHttpTransport(); + }; + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel; + rotateCertificate(channel); + + Thread refreshThread = new Thread(channel::refresh); + Thread shutdownThread = new Thread(channel::shutdown); + try { + refreshThread.start(); + assertTrue(refreshStarted.await(5, TimeUnit.SECONDS)); + shutdownThread.start(); + + shutdownThread.join(200); + assertTrue(shutdownThread.isAlive()); + assertFalse(underlyingChannel.isShutdown()); + } finally { + releaseRefresh.countDown(); + } + refreshThread.join(5000); + shutdownThread.join(5000); + assertFalse(refreshThread.isAlive()); + assertFalse(shutdownThread.isAlive()); + assertEquals(1, channel.getGeneration()); + assertTrue(underlyingChannel.isShutdown()); + } + + @Test + void refresh_concurrentRefreshWhileWaitingForLock_skipsWithoutReadingDisk() throws Exception { + CountDownLatch refreshStarted = new CountDownLatch(1); + CountDownLatch releaseRefresh = new CountDownLatch(1); + transportFactory = + () -> { + if (transportFactoryCount.incrementAndGet() == 2) { + refreshStarted.countDown(); + try { + releaseRefresh.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + return new MockHttpTransport(); + }; + RefreshingHttpJsonChannel channel = createTestChannel(); + rotateCertificate(channel); + + Thread refreshingThread = new Thread(channel::refresh); + Thread waitingThread = new Thread(channel::refresh); + try { + refreshingThread.start(); + assertTrue(refreshStarted.await(5, TimeUnit.SECONDS)); + waitingThread.start(); + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5); + while (waitingThread.getState() != Thread.State.BLOCKED) { + assertTrue(System.nanoTime() < deadline); + Thread.sleep(1); + } + // Rotate again before the first refresh completes. A waiter that re-read the certificate + // would see this fingerprint and refresh a second time. + testFingerprint = "fingerprint3"; + } finally { + releaseRefresh.countDown(); + } + refreshingThread.join(5000); + waitingThread.join(5000); + assertFalse(refreshingThread.isAlive()); + assertFalse(waitingThread.isAlive()); + + // Only the initial transport and the concurrent refresh created transports; the waiter skipped. + assertEquals(2, transportFactoryCount.get()); + assertEquals(1, channel.getGeneration()); + // The second rotation is still detected for the next UNAUTHENTICATED failure. + channel.invalidateDiskFingerprintCache(); + assertTrue(channel.shouldRefresh()); + } + + @Test + void testRefreshFactoryExceptionDoesNotWedgeFingerprint() { + RefreshingHttpJsonChannel channel = createTestChannel(); + HttpTransport initialTransport = channel.getHttpTransport(); + assertEquals(1, transportFactoryCount.get()); + + shouldThrowOnFactory = true; + rotateCertificate(channel); + + // Factory failure is logged and the existing transport is kept + channel.refresh(); + assertEquals(1, transportFactoryCount.get()); + assertEquals(0, channel.getGeneration()); + assertSame(initialTransport, channel.getHttpTransport()); + + // Because the factory threw, the new fingerprint is not recorded as active, so the channel + // still reports that it should be refreshed. + assertTrue(channel.shouldRefresh()); + + shouldThrowOnFactory = false; + channel.refresh(); + assertEquals(2, transportFactoryCount.get()); + assertFalse(channel.shouldRefresh()); + } + + @Test + void testShutdownNowSetsIsShutdown() { + RefreshingHttpJsonChannel channel = createTestChannel(); + assertFalse(channel.isShutdown()); + + channel.shutdownNow(); + + assertTrue(channel.isShutdown()); + } + + @Test + void testAwaitTerminationZeroTimeoutOnTerminatedChannelReturnsTrue() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + lastCreatedChannel.isTerminated = true; + + channel.shutdown(); + assertTrue(channel.awaitTermination(0, TimeUnit.MILLISECONDS)); + } + + @Test + void testChannelDelegationMethods() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel; + FakeHttpJsonClientCall fakeCall = new FakeHttpJsonClientCall<>(); + underlyingChannel.nextCall = fakeCall; + + assertEquals(underlyingChannel.getEndpoint(), channel.getEndpoint()); + assertEquals(underlyingChannel.getHttpTransport(), channel.getHttpTransport()); + assertEquals(underlyingChannel.getExecutor(), channel.getExecutor()); + assertSame(fakeCall, channel.newCall(null, null)); + } + + @Test + void close_shutsDownUnderlyingChannel() { + RefreshingHttpJsonChannel channel = createTestChannel(); + + channel.close(); + + assertTrue(lastCreatedChannel.isShutdown()); + assertTrue(channel.isShutdown()); + } + + @Test + void testConcurrentNewCallDuringRefresh() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + int threadCount = 10; + java.util.concurrent.ExecutorService executorService = + java.util.concurrent.Executors.newFixedThreadPool(threadCount); + java.util.concurrent.CountDownLatch start = new java.util.concurrent.CountDownLatch(1); + java.util.concurrent.CountDownLatch latch = + new java.util.concurrent.CountDownLatch(threadCount); + AtomicInteger successCount = new AtomicInteger(0); + + for (int i = 0; i < threadCount; i++) { + executorService.submit( + () -> { + try { + start.await(); + channel.newCall(null, null); + successCount.incrementAndGet(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } finally { + latch.countDown(); + } + }); + } + + rotateCertificate(channel); + // Release the workers just before refreshing so their calls overlap the transport swap. + start.countDown(); + channel.refresh(); + + assertTrue(latch.await(5, TimeUnit.SECONDS)); + executorService.shutdown(); + assertTrue(executorService.awaitTermination(5, TimeUnit.SECONDS)); + + assertEquals(threadCount, successCount.get()); + assertEquals(1, channel.getGeneration()); + } + + @Test + void testGenerationIncrementAndLifecycleOnDelegatingWrapper() throws Exception { + RefreshingHttpJsonChannel channel = createTestChannel(); + assertEquals(0, channel.getGeneration()); + + rotateCertificate(channel); + channel.refresh(); + + assertEquals(1, channel.getGeneration()); + + // Verify lifecycle methods on delegating wrapper do not throw NullPointerException + assertFalse(channel.isShutdown()); + assertFalse(channel.isTerminated()); + channel.shutdown(); + assertTrue(channel.isShutdown()); + channel.shutdownNow(); + assertTrue(channel.isTerminated()); + assertTrue(channel.awaitTermination(1, TimeUnit.SECONDS)); + } +} diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java index d69fd310c2c0..e6729459b775 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java @@ -189,7 +189,7 @@ public final org.threeten.bp.Duration getInitialRpcTimeout() { * connection has been terminated). * *

{@link #getTotalTimeout()} caps how long the logic should keep trying the RPC until it gives - * up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= + * up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= * totalTimeout. * *

If there are no configurations, Retries have the default initial RPC timeout value of {@code @@ -356,7 +356,7 @@ public final Builder setInitialRpcTimeout(org.threeten.bp.Duration initialTimeou * the connection has been terminated). * *

{@link #getTotalTimeout()} caps how long the logic should keep trying the RPC until it - * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= + * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= * totalTimeout. * *

If there are no configurations, Retries have the default initial RPC timeout value of @@ -491,7 +491,7 @@ public final org.threeten.bp.Duration getInitialRpcTimeout() { * the connection has been terminated). * *

{@link #getTotalTimeout()} caps how long the logic should keep trying the RPC until it - * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= + * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= * totalTimeout. * *

If there are no configurations, Retries have the default initial RPC timeout value of diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java index e4d5461a6b4e..2ad0f6180e6c 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java @@ -102,13 +102,26 @@ public StreamingRetryAlgorithm( (ServerStreamingAttemptException) previousThrowable; previousThrowable = previousThrowable.getCause(); - // If we have made progress in the last attempt, then reset the delays + // If we have made progress in the last attempt, then reset the delays and attempt counts. + // The next attempt is computed from a fresh baseline, so that result algorithms comparing + // the attempt count with the overall attempt count see the reset stream as a new sequence + // of attempts. The previous overall attempt count is then added back so that it keeps + // increasing across resets. if (attemptException.hasSeenResponses()) { - previousSettings = + int previousOverallAttemptCount = previousSettings.getOverallAttemptCount(); + TimedAttemptSettings resetSettings = createFirstAttempt(context).toBuilder() .setFirstAttemptStartTimeNanos(previousSettings.getFirstAttemptStartTimeNanos()) - .setOverallAttemptCount(previousSettings.getOverallAttemptCount()) .build(); + TimedAttemptSettings nextSettings = + super.createNextAttempt(context, previousThrowable, previousResponse, resetSettings); + if (nextSettings == null) { + return null; + } + return nextSettings.toBuilder() + .setOverallAttemptCount( + nextSettings.getOverallAttemptCount() + previousOverallAttemptCount) + .build(); } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java index 67b5e5b285d1..36d6fd966e55 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java @@ -42,7 +42,6 @@ import java.util.Map; import java.util.Set; import javax.annotation.Nonnull; -import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; /** @@ -55,7 +54,6 @@ * *

This is transport specific and each transport has an implementation with its own options. */ -@NullMarked @InternalExtensionOnly public interface ApiCallContext extends RetryingContext { @@ -65,6 +63,18 @@ public interface ApiCallContext extends RetryingContext { /** Returns a new ApiCallContext with the given channel set. */ ApiCallContext withTransportChannel(TransportChannel channel); + /** + * Returns the {@link TransportChannel} associated with this call context, or {@code null} if none + * is set. + * + *

Note: By default, this method returns {@code null}. If an implementation does not override + * this method, automatic mTLS certificate rotation and channel refreshing in retrying callables + * will be disabled. + */ + default TransportChannel getTransportChannel() { + return null; + } + /** Returns a new ApiCallContext with the given Endpoint Context. */ ApiCallContext withEndpointContext(EndpointContext endpointContext); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java index 7d04d38d2605..c8fdbe15f9ee 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java @@ -30,16 +30,67 @@ package com.google.api.gax.rpc; import com.google.api.gax.retrying.BasicResultRetryAlgorithm; +import com.google.api.gax.retrying.RetrySettings; import com.google.api.gax.retrying.RetryingContext; +import com.google.api.gax.retrying.TimedAttemptSettings; import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; /* Package-private for internal use. */ @NullMarked class ApiResultRetryAlgorithm extends BasicResultRetryAlgorithm { - /** Returns true if previousThrowable is an {@link ApiException} that is retryable. */ @Override - public boolean shouldRetry(Throwable previousThrowable, ResponseT previousResponse) { + public @Nullable TimedAttemptSettings createNextAttempt( + @Nullable Throwable previousThrowable, + @Nullable ResponseT previousResponse, + TimedAttemptSettings previousSettings) { + return createNextAttempt(null, previousThrowable, previousResponse, previousSettings); + } + + @Override + public @Nullable TimedAttemptSettings createNextAttempt( + @Nullable RetryingContext context, + @Nullable Throwable previousThrowable, + @Nullable ResponseT previousResponse, + TimedAttemptSettings previousSettings) { + if (isChannelRefreshed(previousThrowable) + && previousSettings.getOverallAttemptCount() == previousSettings.getAttemptCount()) { + RetrySettings globalSettings = previousSettings.getGlobalSettings(); + if (globalSettings.getMaxAttempts() == 0 + && globalSettings.getTotalTimeoutDuration().isZero()) { + globalSettings = globalSettings.toBuilder().setMaxAttempts(1).build(); + } + return previousSettings.toBuilder() + .setGlobalSettings(globalSettings) + .setRetryDelayDuration(java.time.Duration.ZERO) + .setRandomizedRetryDelayDuration(java.time.Duration.ZERO) + .setAttemptCount(previousSettings.getAttemptCount()) + .setOverallAttemptCount(previousSettings.getOverallAttemptCount() + 1) + .build(); + } + if (isChannelRefreshed(previousThrowable)) { + // The single rotation retry has already been used. Return exhausted settings so the retry + // framework stops, rather than returning null and falling back to exponential backoff. + int exhaustedAttemptCount = previousSettings.getAttemptCount() + 1; + return previousSettings.toBuilder() + .setGlobalSettings( + previousSettings.getGlobalSettings().toBuilder() + .setMaxAttempts(exhaustedAttemptCount) + .build()) + .setAttemptCount(exhaustedAttemptCount) + .setOverallAttemptCount(previousSettings.getOverallAttemptCount() + 1) + .build(); + } + return null; + } + + @Override + public boolean shouldRetry( + @Nullable Throwable previousThrowable, @Nullable ResponseT previousResponse) { + if (isChannelRefreshed(previousThrowable)) { + return true; + } return (previousThrowable instanceof ApiException) && ((ApiException) previousThrowable).isRetryable(); } @@ -52,7 +103,15 @@ public boolean shouldRetry(Throwable previousThrowable, ResponseT previousRespon */ @Override public boolean shouldRetry( - RetryingContext context, Throwable previousThrowable, ResponseT previousResponse) { + RetryingContext context, + @Nullable Throwable previousThrowable, + @Nullable ResponseT previousResponse) { + // A failure on a channel that has since been refreshed (e.g. after an mTLS certificate + // rotation) is eligible for a retry regardless of the configured retryable codes; + // createNextAttempt limits it to a single retry. + if (isChannelRefreshed(previousThrowable)) { + return true; + } if (context.getRetryableCodes() != null) { // Ignore the isRetryable() value of the throwable if the RetryingContext has a specific list // of codes that should be retried. @@ -63,4 +122,9 @@ public boolean shouldRetry( } return shouldRetry(previousThrowable, previousResponse); } + + private static boolean isChannelRefreshed(@Nullable Throwable throwable) { + return throwable instanceof UnauthenticatedException + && ((UnauthenticatedException) throwable).isChannelRefreshed(); + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java index 897e1f2dae8e..a64f9f2920fc 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java @@ -35,6 +35,8 @@ import com.google.api.gax.retrying.RetryingFuture; import com.google.common.base.Preconditions; import java.util.concurrent.Callable; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; /** @@ -48,6 +50,7 @@ */ @NullMarked class AttemptCallable implements Callable { + private static final Logger LOG = Logger.getLogger(AttemptCallable.class.getName()); private final UnaryCallable callable; private final RequestT request; private final ApiCallContext originalCallContext; @@ -85,8 +88,45 @@ public ResponseT call() { .getTracer() .attemptStarted(request, externalFuture.getAttemptSettings().getOverallAttemptCount()); + TransportChannel transportChannel = callContext.getTransportChannel(); + final long attemptGeneration = + transportChannel != null ? transportChannel.getGeneration() : 0; + ApiFuture internalFuture = callable.futureCall(request, callContext); - externalFuture.setAttemptFuture(internalFuture); + final ApiCallContext finalContext = callContext; + ApiFuture mappedFuture = + ApiFutures.catching( + internalFuture, + UnauthenticatedException.class, + unauthenticatedException -> { + TransportChannel channel = finalContext.getTransportChannel(); + if (channel != null) { + // If another request already refreshed the channel, retry without checking the + // certificate on disk again. Otherwise, check for a rotation and refresh. + boolean shouldRetry = channel.getGeneration() > attemptGeneration; + if (!shouldRetry) { + try { + if (channel.shouldRefresh()) { + channel.refresh(); + } + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); + } + shouldRetry = channel.getGeneration() > attemptGeneration; + } + + if (shouldRetry) { + throw unauthenticatedException.withChannelRefreshed(); + } + } + throw unauthenticatedException; + }, + com.google.common.util.concurrent.MoreExecutors.directExecutor()); + + externalFuture.setAttemptFuture(mappedFuture); } catch (Throwable e) { externalFuture.setAttemptFuture(ApiFutures.immediateFailedFuture(e)); } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java index 97e22d6ee41f..c8fe6b2a59d3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java @@ -29,6 +29,8 @@ */ package com.google.api.gax.rpc; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -42,6 +44,7 @@ */ @NullMarked public abstract class BidiStreamingCallable { + private static final Logger LOG = Logger.getLogger(BidiStreamingCallable.class.getName()); protected BidiStreamingCallable() {} @@ -241,11 +244,47 @@ public BidiStreamingCallable withDefaultCallContext( return new BidiStreamingCallable() { @Override public ClientStream internalCall( - ResponseObserver responseObserver, + final ResponseObserver responseObserver, ClientStreamReadyObserver onReady, ApiCallContext thisCallContext) { - return BidiStreamingCallable.this.internalCall( - responseObserver, onReady, defaultCallContext.merge(thisCallContext)); + final ApiCallContext mergedContext = defaultCallContext.merge(thisCallContext); + ResponseObserver refreshingObserver = + new ResponseObserver() { + @Override + public void onStart(StreamController controller) { + responseObserver.onStart(controller); + } + + @Override + public void onResponse(ResponseT response) { + responseObserver.onResponse(response); + } + + @Override + public void onError(Throwable t) { + if (t instanceof UnauthenticatedException) { + TransportChannel transportChannel = mergedContext.getTransportChannel(); + if (transportChannel != null && transportChannel.shouldRefresh()) { + try { + transportChannel.refresh(); + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); + } + } + } + responseObserver.onError(t); + } + + @Override + public void onComplete() { + responseObserver.onComplete(); + } + }; + + return BidiStreamingCallable.this.internalCall(refreshingObserver, onReady, mergedContext); } }; } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java index 854ce3fd5de8..5c87e4c558b7 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java @@ -29,6 +29,8 @@ */ package com.google.api.gax.rpc; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -42,6 +44,7 @@ */ @NullMarked public abstract class ClientStreamingCallable { + private static final Logger LOG = Logger.getLogger(ClientStreamingCallable.class.getName()); protected ClientStreamingCallable() {} @@ -77,9 +80,40 @@ public ClientStreamingCallable withDefaultCallContext( return new ClientStreamingCallable() { @Override public ApiStreamObserver clientStreamingCall( - ApiStreamObserver responseObserver, ApiCallContext thisCallContext) { - return ClientStreamingCallable.this.clientStreamingCall( - responseObserver, defaultCallContext.merge(thisCallContext)); + final ApiStreamObserver responseObserver, ApiCallContext thisCallContext) { + final ApiCallContext mergedContext = defaultCallContext.merge(thisCallContext); + ApiStreamObserver refreshingObserver = + new ApiStreamObserver() { + @Override + public void onNext(ResponseT response) { + responseObserver.onNext(response); + } + + @Override + public void onError(Throwable t) { + if (t instanceof UnauthenticatedException) { + TransportChannel transportChannel = mergedContext.getTransportChannel(); + if (transportChannel != null && transportChannel.shouldRefresh()) { + try { + transportChannel.refresh(); + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); + } + } + } + responseObserver.onError(t); + } + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + + return ClientStreamingCallable.this.clientStreamingCall(refreshingObserver, mergedContext); } }; } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java index af6e8014e38a..efc374a50ec3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java @@ -37,6 +37,8 @@ import com.google.errorprone.annotations.concurrent.GuardedBy; import java.util.concurrent.Callable; import java.util.concurrent.CancellationException; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; /** @@ -96,6 +98,8 @@ */ @NullMarked final class ServerStreamingAttemptCallable implements Callable { + private static final Logger LOG = + Logger.getLogger(ServerStreamingAttemptCallable.class.getName()); private final Object lock = new Object(); private final ServerStreamingCallable innerCallable; @@ -221,6 +225,10 @@ public Void call() { .getTracer() .attemptStarted(request, outerRetryingFuture.getAttemptSettings().getOverallAttemptCount()); + final ApiCallContext finalContext = attemptContext; + TransportChannel channelForAttempt = finalContext.getTransportChannel(); + final long attemptGeneration = + channelForAttempt != null ? channelForAttempt.getGeneration() : 0; innerCallable.call( request, new StateCheckingResponseObserver() { @@ -236,6 +244,35 @@ public void onResponseImpl(ResponseT response) { @Override public void onErrorImpl(Throwable t) { + Throwable cause = t; + if (cause instanceof ServerStreamingAttemptException) { + cause = cause.getCause(); + } + if (cause instanceof UnauthenticatedException) { + UnauthenticatedException unauthenticatedException = (UnauthenticatedException) cause; + TransportChannel transportChannel = finalContext.getTransportChannel(); + if (transportChannel != null) { + // If another request already refreshed the channel, retry without checking the + // certificate on disk again. Otherwise, check for a rotation and refresh. + boolean shouldRetry = transportChannel.getGeneration() > attemptGeneration; + if (!shouldRetry) { + try { + if (transportChannel.shouldRefresh()) { + transportChannel.refresh(); + } + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); + } + shouldRetry = transportChannel.getGeneration() > attemptGeneration; + } + if (shouldRetry) { + t = unauthenticatedException.withChannelRefreshed(); + } + } + } onAttemptError(t); } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java index 1866092e28f9..cad80d4f87b3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java @@ -31,10 +31,8 @@ import com.google.api.core.InternalExtensionOnly; import com.google.api.gax.core.BackgroundResource; -import org.jspecify.annotations.NullMarked; /** Class whose instances can issue RPCs on a particular transport. */ -@NullMarked @InternalExtensionOnly public interface TransportChannel extends BackgroundResource { @@ -49,4 +47,28 @@ public interface TransportChannel extends BackgroundResource { * Returns an empty {@link ApiCallContext} that is compatible with this {@code TransportChannel}. */ ApiCallContext getEmptyCallContext(); + + /** + * Refreshes or recreates the underlying network connections of this transport channel. + * + *

By default, this is a no-op for transports that do not require stateful connection lifecycle + * management. + */ + default void refresh() {} + + /** + * Returns true if a certificate rotation has been detected on disk and the transport channel + * should be refreshed, or false otherwise. + */ + default boolean shouldRefresh() { + return false; + } + + /** + * Returns a monotonic generation counter tracking the number of successful refreshes or channel + * rotations performed by this transport channel. + */ + default long getGeneration() { + return 0; + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java index 0c93f07e6b51..d18989264572 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java @@ -30,6 +30,7 @@ package com.google.api.gax.rpc; import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; /** * Exception thrown when the request does not have valid authentication credentials for the @@ -37,18 +38,33 @@ */ @NullMarked public class UnauthenticatedException extends ApiException { + // Pinned to the value computed for previous releases (gax 2.83.0 to 2.87.0) so that adding + // members does not break Java serialization compatibility with them. + private static final long serialVersionUID = 6971115068105015909L; + + /** + * Whether this failure happened on a transport channel that has since been refreshed (for + * example, after an mTLS certificate rotation), making the request eligible for a single + * immediate retry on the refreshed channel. Only meaningful within the process that observed the + * failure, so it is not serialized. + */ + private final transient boolean channelRefreshed; + public UnauthenticatedException(Throwable cause, StatusCode statusCode, boolean retryable) { super(cause, statusCode, retryable); + this.channelRefreshed = false; } public UnauthenticatedException( String message, Throwable cause, StatusCode statusCode, boolean retryable) { super(message, cause, statusCode, retryable); + this.channelRefreshed = false; } public UnauthenticatedException( Throwable cause, StatusCode statusCode, boolean retryable, ErrorDetails errorDetails) { super(cause, statusCode, retryable, errorDetails); + this.channelRefreshed = false; } public UnauthenticatedException( @@ -58,5 +74,38 @@ public UnauthenticatedException( boolean retryable, ErrorDetails errorDetails) { super(message, cause, statusCode, retryable, errorDetails); + this.channelRefreshed = false; + } + + private UnauthenticatedException( + @Nullable String message, + @Nullable Throwable cause, + StatusCode statusCode, + boolean retryable, + @Nullable ErrorDetails errorDetails, + boolean channelRefreshed) { + super(message, cause, statusCode, retryable, errorDetails); + this.channelRefreshed = channelRefreshed; + } + + /** Returns whether this failure happened on a transport channel that has since been refreshed. */ + boolean isChannelRefreshed() { + return channelRefreshed; + } + + /** + * Returns a copy of this exception marked as having happened on a channel that has since been + * refreshed. The copy keeps the message, cause, status code, {@link #isRetryable()} value, error + * details, stack trace and suppressed exceptions of this exception. + */ + UnauthenticatedException withChannelRefreshed() { + UnauthenticatedException newEx = + new UnauthenticatedException( + getMessage(), getCause(), getStatusCode(), isRetryable(), getErrorDetails(), true); + newEx.setStackTrace(getStackTrace()); + for (Throwable suppressed : getSuppressed()) { + newEx.addSuppressed(suppressed); + } + return newEx; } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index 99f2cca9b53d..d77378ad9612 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -32,21 +32,20 @@ import com.google.api.core.InternalApi; import com.google.api.gax.rpc.internal.EnvironmentProvider; -import org.jspecify.annotations.NullMarked; +import com.google.auth.mtls.MtlsUtils; +import com.google.auth.oauth2.PropertyProvider; /** * Utility class for handling certificate-based access configurations. * - *

This class handles the processing of GOOGLE_API_USE_CLIENT_CERTIFICATE and - * GOOGLE_API_USE_MTLS_ENDPOINT environment variables according to https://google.aip.dev/auth/4114 + *

This class handles the processing of GOOGLE_API_USE_CLIENT_CERTIFICATE, + * GOOGLE_API_CERTIFICATE_CONFIG, and GOOGLE_API_USE_MTLS_ENDPOINT configurations. */ -@NullMarked @InternalApi public class CertificateBasedAccess { private final EnvironmentProvider envProvider; - /** The EnvironmentProvider mechanism supports env var injection for unit tests. */ public CertificateBasedAccess(EnvironmentProvider envProvider) { this.envProvider = envProvider; } @@ -66,20 +65,35 @@ public enum MtlsEndpointUsagePolicy { ALWAYS; } + private com.google.auth.oauth2.EnvironmentProvider getAuthEnvProvider() { + return name -> envProvider.getenv(name); + } + + private PropertyProvider getAuthPropertyProvider() { + return System::getProperty; + } + /** Returns if mutual TLS client certificate should be used. */ public boolean useMtlsClientCertificate() { - String useClientCertificate = envProvider.getenv("GOOGLE_API_USE_CLIENT_CERTIFICATE"); - return "true".equals(useClientCertificate); + return MtlsUtils.useMtlsClientCertificate(getAuthEnvProvider(), getAuthPropertyProvider()); } /** Returns the current mutual TLS endpoint usage policy. */ public MtlsEndpointUsagePolicy getMtlsEndpointUsagePolicy() { String mtlsEndpointUsagePolicy = envProvider.getenv("GOOGLE_API_USE_MTLS_ENDPOINT"); - if ("never".equals(mtlsEndpointUsagePolicy)) { + if ("never".equalsIgnoreCase(mtlsEndpointUsagePolicy)) { return MtlsEndpointUsagePolicy.NEVER; - } else if ("always".equals(mtlsEndpointUsagePolicy)) { + } else if ("always".equalsIgnoreCase(mtlsEndpointUsagePolicy)) { return MtlsEndpointUsagePolicy.ALWAYS; } return MtlsEndpointUsagePolicy.AUTO; } + + /** + * Resolves and returns the path to the mutual TLS client certificate, or null if none should be + * used. + */ + public String getWorkloadCertPath() { + return MtlsUtils.getWorkloadCertPath(getAuthEnvProvider(), getAuthPropertyProvider()); + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateRotationTracker.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateRotationTracker.java new file mode 100644 index 000000000000..fc58d67ab1d6 --- /dev/null +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateRotationTracker.java @@ -0,0 +1,197 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.rpc.mtls; + +import com.google.api.core.InternalApi; +import com.google.common.base.Strings; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.locks.ReentrantLock; +import java.util.function.Function; +import java.util.function.Supplier; +import org.jspecify.annotations.Nullable; + +/** + * Thread-safe helper that tracks workload certificate fingerprints on disk to detect mTLS + * certificate rotations for transport channels ({@code ChannelPool} and {@code + * RefreshingHttpJsonChannel}). + * + *

For internal use only. + */ +@InternalApi +public class CertificateRotationTracker { + + /** + * Duration (1 second) for which a detected certificate rotation (where the disk fingerprint + * differs from {@code activeCertFingerprint}) is cached in memory. + * + *

When a certificate rotates on disk, many concurrent in-flight RPCs may fail with {@code + * UNAUTHENTICATED} simultaneously and call {@link #shouldRefresh()}. Caching positive rotation + * detections for 1 second coalesces disk reads ({@code Files.readAllBytes}) and SHA-256 hashing + * across concurrent threads to prevent a thundering herd of file I/O, while bounding maximum + * staleness to 1 second. + */ + private static final long POSITIVE_ROTATION_CACHE_TTL_NANOS = TimeUnit.SECONDS.toNanos(1); + + private static final class DiskCheckResult { + final String fingerprint; + final long sequence; + final long timestampNanos; + + DiskCheckResult(String fingerprint, long sequence, long timestampNanos) { + this.fingerprint = fingerprint; + this.sequence = sequence; + this.timestampNanos = timestampNanos; + } + } + + private final Supplier certPathSupplier; + private final Function fingerprintReader; + private volatile String activeCertFingerprint; + private volatile DiskCheckResult lastDiskCheck = null; + private final ReentrantLock diskCheckLock = new ReentrantLock(); + private final AtomicLong diskCheckSequence = new AtomicLong(0); + + /** + * Creates a tracker for a fixed workload certificate path using {@link + * WorkloadCertificateUtils#getCertificateFingerprint(String)}. + */ + public CertificateRotationTracker(@Nullable String workloadCertPath) { + this(() -> workloadCertPath, WorkloadCertificateUtils::getCertificateFingerprint); + } + + /** + * Creates a tracker with custom certificate path and fingerprint suppliers (used by transports + * that expose package-private overrides for testing). + */ + public CertificateRotationTracker( + Supplier certPathSupplier, Function fingerprintReader) { + this.certPathSupplier = certPathSupplier; + this.fingerprintReader = fingerprintReader; + String initialCertPath = certPathSupplier.get(); + this.activeCertFingerprint = + initialCertPath != null + ? Strings.nullToEmpty(fingerprintReader.apply(initialCertPath)) + : ""; + } + + /** + * Returns {@code true} if a workload certificate path is configured, readable, and its current + * SHA-256 fingerprint on disk differs from the active fingerprint. + */ + public boolean shouldRefresh() { + String certPath = certPathSupplier.get(); + if (certPath == null) { + return false; + } + String currentDiskFingerprint = getOrUpdateDiskFingerprint(certPath); + if (currentDiskFingerprint.isEmpty()) { + return false; + } + return !currentDiskFingerprint.equalsIgnoreCase(activeCertFingerprint); + } + + private String getOrUpdateDiskFingerprint(String certPath) { + long seqBeforeLock = diskCheckSequence.get(); + long now = System.nanoTime(); + DiskCheckResult cached = lastDiskCheck; + if (cached != null + && !cached.fingerprint.isEmpty() + && !cached.fingerprint.equalsIgnoreCase(this.activeCertFingerprint) + && (now - cached.timestampNanos < POSITIVE_ROTATION_CACHE_TTL_NANOS)) { + return cached.fingerprint; + } + + diskCheckLock.lock(); + try { + now = System.nanoTime(); + cached = lastDiskCheck; + if (cached != null + && !cached.fingerprint.isEmpty() + && (cached.sequence > seqBeforeLock + || (!cached.fingerprint.equalsIgnoreCase(this.activeCertFingerprint) + && (now - cached.timestampNanos < POSITIVE_ROTATION_CACHE_TTL_NANOS)))) { + return cached.fingerprint; + } + long newSeq = diskCheckSequence.incrementAndGet(); + String fingerprint = Strings.nullToEmpty(fingerprintReader.apply(certPath)); + if (!fingerprint.isEmpty()) { + lastDiskCheck = new DiskCheckResult(fingerprint, newSeq, System.nanoTime()); + } else { + lastDiskCheck = null; + } + return fingerprint; + } finally { + diskCheckLock.unlock(); + } + } + + /** + * Reads the current certificate fingerprint directly from disk (bypassing the 1-second cache), + * returning {@code ""} if no workload certificate path is configured or if the file is currently + * unreadable/empty. + */ + public String readDiskFingerprint() { + String certPath = certPathSupplier.get(); + if (certPath == null) { + return ""; + } + return Strings.nullToEmpty(fingerprintReader.apply(certPath)); + } + + /** + * Returns {@code true} if {@code diskFingerprint} matches the currently active certificate + * fingerprint. + */ + public boolean isAlreadyActive(String diskFingerprint) { + return diskFingerprint != null && diskFingerprint.equalsIgnoreCase(this.activeCertFingerprint); + } + + /** + * Updates the active certificate fingerprint after a successful channel refresh and clears any + * cached disk check result. + */ + public void markRefreshed(String newFingerprint) { + if (newFingerprint != null && !newFingerprint.isEmpty()) { + this.activeCertFingerprint = newFingerprint; + this.lastDiskCheck = null; + } + } + + /** Returns the currently active certificate fingerprint (or {@code ""} if none). */ + public String getActiveCertFingerprint() { + return activeCertFingerprint; + } + + /** Invalidates the cached disk check result. Visible for testing. */ + public void invalidateCache() { + this.lastDiskCheck = null; + } +} diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java new file mode 100644 index 000000000000..94fb36db1ac3 --- /dev/null +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java @@ -0,0 +1,64 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.rpc.mtls; + +import com.google.api.core.InternalApi; +import com.google.auth.mtls.MtlsUtils; +import java.io.File; + +/** Internal utility class for managing dynamic workload certificates. */ +@InternalApi +public class WorkloadCertificateUtils { + + private static final String EMPTY_FILE_SHA256 = + "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"; + + private WorkloadCertificateUtils() {} + + /** + * Computes the SHA-256 fingerprint of the certificate file at {@code certPath}, returning {@code + * ""} if the path is {@code null}, unreadable, empty (e.g., temporarily truncated to 0 bytes + * mid-write by an external certificate rotator), or hashes to the empty-byte digest. + * + *

Returning {@code ""} on unreadable or empty files ensures callers ({@code shouldRefresh()} + * and {@code refresh()}) safely skip refreshing during transient mid-write states rather than + * treating an empty digest as a certificate rotation mismatch. + */ + public static String getCertificateFingerprint(String certPath) { + if (certPath == null || new File(certPath).length() == 0) { + return ""; + } + String fingerprint = MtlsUtils.getCertificateFingerprint(certPath); + if (fingerprint == null || EMPTY_FILE_SHA256.equalsIgnoreCase(fingerprint)) { + return ""; + } + return fingerprint; + } +} diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java index 300a0ad30130..1e730d2db1f3 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java @@ -29,14 +29,25 @@ */ package com.google.api.gax.rpc; +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.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import com.google.api.core.NanoClock; +import com.google.api.gax.retrying.ExponentialRetryAlgorithm; +import com.google.api.gax.retrying.RetryAlgorithm; +import com.google.api.gax.retrying.RetrySettings; +import com.google.api.gax.retrying.ServerStreamingAttemptException; +import com.google.api.gax.retrying.StreamingRetryAlgorithm; +import com.google.api.gax.retrying.TimedAttemptSettings; import com.google.api.gax.rpc.StatusCode.Code; import com.google.api.gax.rpc.testing.FakeStatusCode; import com.google.common.collect.Sets; +import java.time.Duration; import java.util.Collections; import org.junit.jupiter.api.Test; import org.mockito.Mockito; @@ -111,4 +122,477 @@ void testShouldRetryWithContextWithEmptyRetryableCodes() { ApiResultRetryAlgorithm algorithm = new ApiResultRetryAlgorithm<>(); assertFalse(algorithm.shouldRetry(context, unavailableException, null)); } + + @Test + void testRotationRetryWithNonRetryableSettings_maxAttemptsOne() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Collections.emptySet()); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(1) + .setInitialRpcTimeoutDuration(Duration.ofSeconds(10)) + .setMaxRpcTimeoutDuration(Duration.ofSeconds(10)) + .setTotalTimeoutDuration(Duration.ofSeconds(10)) + .build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = rotationException(); + + // First rotation failure: grants immediate free retry without incrementing attemptCount + TimedAttemptSettings nextAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertNotNull(nextAttempt); + assertEquals(Duration.ZERO, nextAttempt.getRetryDelayDuration()); + assertEquals(Duration.ZERO, nextAttempt.getRandomizedRetryDelayDuration()); + assertEquals(0, nextAttempt.getAttemptCount()); + assertEquals(1, nextAttempt.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, nextAttempt)); + + // Second consecutive failure: overallAttemptCount (1) != attemptCount (0), so no free retry + TimedAttemptSettings thirdAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, nextAttempt); + assertNotNull(thirdAttempt); + assertEquals(1, thirdAttempt.getAttemptCount()); + assertEquals(2, thirdAttempt.getOverallAttemptCount()); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, thirdAttempt)); + } + + @Test + void testRotationRetryWithNonRetryableSettings_zeroMaxAttemptsZeroTotalTimeout() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Collections.emptySet()); + + RetrySettings settings = + RetrySettings.newBuilder().setMaxAttempts(0).setTotalTimeoutDuration(Duration.ZERO).build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = rotationException(); + + TimedAttemptSettings nextAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertNotNull(nextAttempt); + assertEquals(1, nextAttempt.getGlobalSettings().getMaxAttempts()); + assertEquals(0, nextAttempt.getAttemptCount()); + assertEquals(1, nextAttempt.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, nextAttempt)); + + // Subsequent failure is rejected + TimedAttemptSettings thirdAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, nextAttempt); + assertNotNull(thirdAttempt); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, thirdAttempt)); + } + + @Test + void testRotationRetryAfterTransientErrorPreservesRemainingBudget() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(3) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofSeconds(30)) + .build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings attempt0 = retryAlgorithm.createFirstAttempt(context); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); + + // Attempt 0 fails with UNAVAILABLE -> normal retry (attemptCount = 1, overallAttemptCount = 1) + TimedAttemptSettings attempt1 = + retryAlgorithm.createNextAttempt(context, unavailableEx, null, attempt0); + assertEquals(1, attempt1.getAttemptCount()); + assertEquals(1, attempt1.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, unavailableEx, null, attempt1)); + + // Attempt 1 fails with rotation 401 -> free zero-delay retry (attemptCount = 1, + // overallAttemptCount = 2) + TimedAttemptSettings attempt2 = + retryAlgorithm.createNextAttempt(context, rotationEx, null, attempt1); + assertEquals(Duration.ZERO, attempt2.getRetryDelayDuration()); + assertEquals(1, attempt2.getAttemptCount()); + assertEquals(2, attempt2.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, attempt2)); + + // Attempt 2 fails with UNAVAILABLE -> normal retry still allowed (attemptCount = 2 < + // maxAttempts = 3) + TimedAttemptSettings attempt3 = + retryAlgorithm.createNextAttempt(context, unavailableEx, null, attempt2); + assertEquals(2, attempt3.getAttemptCount()); + assertEquals(3, attempt3.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, unavailableEx, null, attempt3)); + } + + @Test + void testSecondRotationFailureStopsWithTotalTimeoutAndNoMaxAttempts() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Collections.emptySet()); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(0) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = rotationException(); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, rotationRetry)); + + // A second rotation-marked failure must stop immediately instead of retrying with backoff + // until the total timeout expires. + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(context, rotationEx, null, rotationRetry); + assertNotNull(afterSecondFailure); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, afterSecondFailure)); + } + + @Test + void testSecondRotationFailureStopsDespiteRemainingMaxAttempts() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = rotationException(); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, rotationRetry)); + + // UNAUTHENTICATED is not a retryable code, so a second rotation-marked failure must stop + // instead of consuming the method's remaining attempts. + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(context, rotationEx, null, rotationRetry); + assertNotNull(afterSecondFailure); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, afterSecondFailure)); + } + + @Test + void testStreamRotationRetryAfterProgressNotBlockedByEarlierRetry() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + + StreamingRetryAlgorithm retryAlgorithm = + new StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); + + // The first attempt fails with UNAVAILABLE before receiving any messages: normal retry. + TimedAttemptSettings attempt1 = + retryAlgorithm.createNextAttempt( + context, + new ServerStreamingAttemptException(unavailableEx, true, false), + null, + retryAlgorithm.createFirstAttempt(context)); + assertEquals(1, attempt1.getAttemptCount()); + assertEquals(1, attempt1.getOverallAttemptCount()); + + // The second attempt receives messages and then fails with a rotation error. The earlier + // retry must not prevent the free rotation retry. + ServerStreamingAttemptException rotationAfterProgress = + new ServerStreamingAttemptException(rotationEx, true, true); + TimedAttemptSettings attempt2 = + retryAlgorithm.createNextAttempt(context, rotationAfterProgress, null, attempt1); + assertNotNull(attempt2); + assertEquals(Duration.ZERO, attempt2.getRetryDelayDuration()); + assertEquals(Duration.ZERO, attempt2.getRandomizedRetryDelayDuration()); + assertEquals(0, attempt2.getAttemptCount()); + assertEquals(2, attempt2.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationAfterProgress, null, attempt2)); + + // A repeated rotation error without further progress stops the stream. + ServerStreamingAttemptException rotationWithoutProgress = + new ServerStreamingAttemptException(rotationEx, true, false); + TimedAttemptSettings attempt3 = + retryAlgorithm.createNextAttempt(context, rotationWithoutProgress, null, attempt2); + assertNotNull(attempt3); + assertEquals(3, attempt3.getOverallAttemptCount()); + assertFalse(retryAlgorithm.shouldRetry(context, rotationWithoutProgress, null, attempt3)); + } + + @Test + void testStreamProgressResetKeepsOverallAttemptCountIncreasing() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + + StreamingRetryAlgorithm retryAlgorithm = + new StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + ServerStreamingAttemptException unavailableWithoutProgress = + new ServerStreamingAttemptException(unavailableEx, true, false); + ServerStreamingAttemptException unavailableAfterProgress = + new ServerStreamingAttemptException(unavailableEx, true, true); + + TimedAttemptSettings attempt1 = + retryAlgorithm.createNextAttempt( + context, unavailableWithoutProgress, null, retryAlgorithm.createFirstAttempt(context)); + TimedAttemptSettings attempt2 = + retryAlgorithm.createNextAttempt(context, unavailableWithoutProgress, null, attempt1); + assertEquals(2, attempt2.getAttemptCount()); + assertEquals(2, attempt2.getOverallAttemptCount()); + + // Progress resets the attempt count, but the overall attempt count keeps increasing. + TimedAttemptSettings attempt3 = + retryAlgorithm.createNextAttempt(context, unavailableAfterProgress, null, attempt2); + assertEquals(1, attempt3.getAttemptCount()); + assertEquals(3, attempt3.getOverallAttemptCount()); + assertEquals(settings.getInitialRetryDelayDuration(), attempt3.getRetryDelayDuration()); + assertEquals( + attempt2.getFirstAttemptStartTimeNanos(), attempt3.getFirstAttemptStartTimeNanos()); + assertTrue(retryAlgorithm.shouldRetry(context, unavailableAfterProgress, null, attempt3)); + } + + @Test + void testStreamProgressResetReturnsNullWhenNotRetryable() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Collections.emptySet()); + + StreamingRetryAlgorithm retryAlgorithm = + new StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm( + RetrySettings.newBuilder().setMaxAttempts(5).build(), NanoClock.getDefaultClock())); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + + TimedAttemptSettings next = + retryAlgorithm.createNextAttempt( + context, + new ServerStreamingAttemptException(unavailableEx, true, true), + null, + retryAlgorithm.createFirstAttempt(context)); + assertNull(next); + } + + @Test + void testConfiguredUnauthenticated_notFlagged_usesNormalBackoff() { + // No per-call retryable codes: the method's configuration is carried by isRetryable(). + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(null); + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + // UNAUTHENTICATED configured as retryable, not caused by a channel refresh. + UnauthenticatedException configuredEx = + new UnauthenticatedException( + "Invalid token", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + TimedAttemptSettings attempt = retryAlgorithm.createFirstAttempt(context); + for (int i = 1; i < 5; i++) { + attempt = retryAlgorithm.createNextAttempt(context, configuredEx, null, attempt); + assertNotNull(attempt); + assertEquals(i, attempt.getAttemptCount()); + assertEquals(i, attempt.getOverallAttemptCount()); + assertTrue(attempt.getRetryDelayDuration().compareTo(Duration.ZERO) > 0); + assertTrue(retryAlgorithm.shouldRetry(context, configuredEx, null, attempt)); + } + // The fifth failure exhausts maxAttempts = 5. + attempt = retryAlgorithm.createNextAttempt(context, configuredEx, null, attempt); + assertFalse(retryAlgorithm.shouldRetry(context, configuredEx, null, attempt)); + } + + @Test + void testContextRetryableCodesExcludeUnauthenticated_notFlagged_notRetried() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + // Retryable according to the method's configuration, but the per-call codes exclude it. + UnauthenticatedException configuredEx = + new UnauthenticatedException( + "Invalid token", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + ApiResultRetryAlgorithm algorithm = new ApiResultRetryAlgorithm<>(); + assertFalse(algorithm.shouldRetry(context, configuredEx, null)); + assertNull( + algorithm.createNextAttempt( + context, + configuredEx, + null, + new ExponentialRetryAlgorithm( + RetrySettings.newBuilder().setMaxAttempts(5).build(), + NanoClock.getDefaultClock()) + .createFirstAttempt())); + } + + @Test + void testFlagged_nullContext_retriesOnce() { + RetrySettings settings = RetrySettings.newBuilder().setMaxAttempts(1).build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + UnauthenticatedException rotationEx = rotationException(); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt( + null, rotationEx, null, retryAlgorithm.createFirstAttempt(null)); + assertNotNull(rotationRetry); + assertEquals(Duration.ZERO, rotationRetry.getRetryDelayDuration()); + assertEquals(0, rotationRetry.getAttemptCount()); + assertEquals(1, rotationRetry.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(null, rotationEx, null, rotationRetry)); + + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(null, rotationEx, null, rotationRetry); + assertFalse(retryAlgorithm.shouldRetry(null, rotationEx, null, afterSecondFailure)); + } + + @Test + void testFlagged_contextWithoutRetryableCodes_retriesOnce() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(null); + RetrySettings settings = RetrySettings.newBuilder().setMaxAttempts(1).build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + // Not retryable by configuration (isRetryable() == false), but caused by a channel refresh. + UnauthenticatedException rotationEx = rotationException(); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt( + context, rotationEx, null, retryAlgorithm.createFirstAttempt(context)); + assertNotNull(rotationRetry); + assertEquals(Duration.ZERO, rotationRetry.getRetryDelayDuration()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, rotationRetry)); + + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(context, rotationEx, null, rotationRetry); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, afterSecondFailure)); + } + + @Test + void testConfiguredUnauthenticated_afterRotationRetry_continuesWithNormalPolicy() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(null); + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(3) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + UnauthenticatedException configuredEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = configuredEx.withChannelRefreshed(); + + // A rotation failure gets the free zero-delay retry without consuming an attempt. + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt( + context, rotationEx, null, retryAlgorithm.createFirstAttempt(context)); + assertEquals(Duration.ZERO, rotationRetry.getRetryDelayDuration()); + assertEquals(0, rotationRetry.getAttemptCount()); + + // A later, unrelated failure on the same channel follows the configured policy with backoff. + TimedAttemptSettings normalRetry = + retryAlgorithm.createNextAttempt(context, configuredEx, null, rotationRetry); + assertNotNull(normalRetry); + assertEquals(1, normalRetry.getAttemptCount()); + assertTrue(normalRetry.getRetryDelayDuration().compareTo(Duration.ZERO) > 0); + assertTrue(retryAlgorithm.shouldRetry(context, configuredEx, null, normalRetry)); + } + + /** + * An UNAUTHENTICATED failure caused by a channel refresh, where UNAUTHENTICATED is not configured + * as retryable (the common case). + */ + private static UnauthenticatedException rotationException() { + return new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ false) + .withChannelRefreshed(); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java index 4b5da578862c..a1e148c241b6 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java @@ -30,16 +30,23 @@ package com.google.api.gax.rpc; import static com.google.common.truth.Truth.assertThat; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import com.google.api.core.ApiFuture; import com.google.api.core.SettableApiFuture; import com.google.api.gax.retrying.RetrySettings; import com.google.api.gax.retrying.RetryingFuture; import com.google.api.gax.retrying.TimedAttemptSettings; import com.google.api.gax.rpc.testing.FakeCallContext; +import com.google.api.gax.rpc.testing.FakeChannel; +import com.google.api.gax.rpc.testing.FakeStatusCode; +import com.google.api.gax.rpc.testing.FakeTransportChannel; import com.google.api.gax.tracing.ApiTracer; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; @@ -133,4 +140,388 @@ void testRpcTimeoutIsNotErased() { assertThat(capturedCallContext.getValue().getTimeoutDuration()).isEqualTo(callerTimeout); } + + @Test + void testRefreshedUnauthenticated_flaggedAndPreservesContext() { + FakeTransportChannel transportChannel = + FakeTransportChannel.create(new FakeChannel()).setShouldRefresh(true); + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", + new IllegalStateException("Root cause"), + FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), + false); + originalEx.setStackTrace( + new StackTraceElement[] {new StackTraceElement("foo", "bar", "Baz.java", 123)}); + originalEx.addSuppressed(new RuntimeException("Suppressed cause")); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isChannelRefreshed()).isTrue(); + // The original retryable setting is preserved. + assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.getMessage()).isEqualTo(originalEx.getMessage()); + assertThat(rethrown.getStatusCode()).isEqualTo(originalEx.getStatusCode()); + assertThat(rethrown.getCause()).isEqualTo(originalEx.getCause()); + assertThat(rethrown.getStackTrace()).isEqualTo(originalEx.getStackTrace()); + assertThat(rethrown.getSuppressed().length).isEqualTo(1); + assertThat(rethrown.getSuppressed()[0]).isInstanceOf(RuntimeException.class); + } + + @Test + void testSiblingInFlightRequest_channelRotatedInFlight_flaggedWithoutDuplicateRefresh() { + FakeChannel innerChannel = new FakeChannel(); + innerChannel.setGeneration(1); + // Initially shouldRefresh is false because sibling request already completed the refresh + innerChannel.setShouldRefresh(false); + FakeTransportChannel transportChannel = FakeTransportChannel.create(innerChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())) + .thenAnswer( + invocation -> { + // While request was in flight, sibling finished rotation and bumped generation to 2 + innerChannel.setGeneration(2); + return failedFuture; + }); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + // Sibling request should be flagged for a retry on the new channel + assertThat(rethrown.isChannelRefreshed()).isTrue(); + // But should NOT have triggered a second refresh call + assertThat(innerChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testPermanentUnauthenticatedFailure_sameGeneration_notFlagged() { + FakeChannel innerChannel = new FakeChannel(); + innerChannel.setGeneration(1); + innerChannel.setShouldRefresh(false); + FakeTransportChannel transportChannel = FakeTransportChannel.create(innerChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Invalid credentials", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + // Genuine permanent error on same generation is not flagged for a rotation retry + assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); + assertThat(innerChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testRefreshThrowsException_notFlaggedWhenGenerationUnchanged() { + AtomicInteger refreshCalls = new AtomicInteger(); + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + return true; + } + + @Override + public void refresh() { + refreshCalls.incrementAndGet(); + throw new RuntimeException("Refresh error"); + } + }; + FakeTransportChannel transportChannel = FakeTransportChannel.create(fakeChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); + assertThat(refreshCalls.get()).isEqualTo(1); + } + + @Test + void testRefreshReturnsWithoutAdvancingGeneration_notFlagged() { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + return true; + } + + @Override + public void refresh() { + // Simulates refresh() returning early without rotating any channel (e.g. unreadable + // cert) + } + }; + FakeTransportChannel transportChannel = FakeTransportChannel.create(fakeChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); + } + + @Test + void testRefreshThrowsException_flaggedIfConcurrentThreadAdvancedGeneration() { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + return true; + } + + @Override + public void refresh() { + // Concurrent thread advanced generation before/during refresh failure + setGeneration(getGeneration() + 1); + throw new RuntimeException("Refresh error on this thread"); + } + }; + FakeTransportChannel transportChannel = FakeTransportChannel.create(fakeChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isChannelRefreshed()).isTrue(); + } + + @Test + void testGenerationAlreadyAdvanced_flaggedWithoutCheckingShouldRefresh() { + AtomicInteger shouldRefreshCalls = new AtomicInteger(); + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + shouldRefreshCalls.incrementAndGet(); + return true; + } + }; + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())) + .thenAnswer( + invocation -> { + // Another request rotated the channel while this one was in flight. + fakeChannel.setGeneration(fakeChannel.getGeneration() + 1); + return failedFuture; + }); + + UnauthenticatedException rethrown = callAndGetUnauthenticated(fakeChannel); + + assertThat(rethrown.isChannelRefreshed()).isTrue(); + // The certificate on disk is not checked again, and no second refresh happens. + assertThat(shouldRefreshCalls.get()).isEqualTo(0); + assertThat(fakeChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testShouldRefreshThrows_originalUnauthenticatedPreserved() { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + throw new IllegalStateException("Unable to read certificate"); + } + }; + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + UnauthenticatedException rethrown = callAndGetUnauthenticated(fakeChannel); + + assertThat(rethrown).isSameInstanceAs(originalEx); + assertThat(rethrown.isChannelRefreshed()).isFalse(); + assertThat(fakeChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testConfiguredRetryableUnauthenticated_sameGeneration_notFlagged() { + FakeChannel fakeChannel = new FakeChannel().setShouldRefresh(false); + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Token expired", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), true); + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + UnauthenticatedException rethrown = callAndGetUnauthenticated(fakeChannel); + + // Left to the normal retry policy: still retryable, but not a rotation retry. + assertThat(rethrown).isSameInstanceAs(originalEx); + assertThat(rethrown.isRetryable()).isTrue(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); + } + + private UnauthenticatedException callAndGetUnauthenticated(FakeChannel fakeChannel) { + ApiCallContext callContext = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + ExecutionException e = + assertThrows(ExecutionException.class, () -> futureCaptor.getValue().get()); + assertThat(e.getCause()).isInstanceOf(UnauthenticatedException.class); + return (UnauthenticatedException) e.getCause(); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java index 4c1b5320de4c..c53ca1509430 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java @@ -77,8 +77,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsFalse() throws IOException { FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = false; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -97,8 +99,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_mtlsUsageAuto() throws IOExc FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -117,8 +121,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_mtlsUsageAlways() throws IOE FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -137,8 +143,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_mtlsUsageNever() throws IOEx FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "never" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -156,8 +164,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_useCertificateIsFalse_nullMt MtlsProvider mtlsProvider = new FakeMtlsProvider(null, "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "false"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(false); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -174,8 +184,10 @@ void mtlsEndpointResolver_getKeyStore_throwsIOException() throws IOException { MtlsProvider mtlsProvider = new FakeMtlsProvider(null, "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); assertThrows( IOException.class, () -> @@ -272,8 +284,10 @@ void endpointContextBuild_mtlsConfigured_GDU() throws IOException { MtlsProvider mtlsProvider = new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); EndpointContext endpointContext = defaultEndpointContextBuilder .setClientSettingsEndpoint(null) @@ -293,8 +307,10 @@ void endpointContextBuild_mtlsConfigured_nonGDU_throwsIllegalArgumentException() MtlsProvider mtlsProvider = new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); EndpointContext.Builder endpointContextBuilder = defaultEndpointContextBuilder .setUniverseDomain("random.com") diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java index 4cdd3cc018cc..5d93f3894874 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java @@ -29,6 +29,7 @@ */ package com.google.api.gax.rpc; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.mockito.Mockito.mock; import com.google.api.core.AbstractApiFuture; @@ -41,6 +42,9 @@ import com.google.api.gax.rpc.StatusCode.Code; import com.google.api.gax.rpc.testing.FakeApiException; import com.google.api.gax.rpc.testing.FakeCallContext; +import com.google.api.gax.rpc.testing.FakeChannel; +import com.google.api.gax.rpc.testing.FakeStatusCode; +import com.google.api.gax.rpc.testing.FakeTransportChannel; import com.google.api.gax.rpc.testing.MockStreamingApi.MockServerStreamingCall; import com.google.api.gax.rpc.testing.MockStreamingApi.MockServerStreamingCallable; import com.google.api.gax.tracing.BaseApiTracer; @@ -248,6 +252,208 @@ void testInitialRetry() { Truth.assertThat(call.getRequest()).isEqualTo("request > 0"); } + @Test + @SuppressWarnings("ConstantConditions") + void testUnauthenticatedRefreshWithoutGenerationAdvance_notFlagged() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + + // Send initial error + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + // Should notify the outer future + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(((ServerStreamingAttemptException) outerError).hasSeenResponses()).isFalse(); + Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); + Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isFalse(); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isChannelRefreshed()) + .isFalse(); + } + + @Test + @SuppressWarnings("ConstantConditions") + void testUnauthenticatedRefreshWithGenerationAdvance_flagsChannelRefreshed() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + java.util.concurrent.atomic.AtomicLong generation = + new java.util.concurrent.atomic.AtomicLong(0); + Mockito.when(transportChannel.getGeneration()).thenAnswer(inv -> generation.get()); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + Mockito.doAnswer( + inv -> { + generation.incrementAndGet(); + return null; + }) + .when(transportChannel) + .refresh(); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); + Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isChannelRefreshed()) + .isTrue(); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isFalse(); + Truth.assertThat(outerError.getCause().getStackTrace()).isEqualTo(initialError.getStackTrace()); + // A fixed clock keeps the fake attempt (started at t=0) within its total timeout. + com.google.api.gax.retrying.StreamingRetryAlgorithm retryAlgorithm = + new com.google.api.gax.retrying.StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm<>(), + new com.google.api.gax.retrying.ExponentialRetryAlgorithm( + RetrySettings.newBuilder().build(), new com.google.api.gax.core.FakeApiClock(0))); + // Mirror the retry executor: compute the next attempt, then ask whether to run it. + TimedAttemptSettings nextAttempt = + retryAlgorithm.createNextAttempt(outerError, null, fakeRetryingFuture.getAttemptSettings()); + Truth.assertThat(nextAttempt).isNotNull(); + Truth.assertThat(retryAlgorithm.shouldRetry(outerError, null, nextAttempt)).isTrue(); + + // Verify retry call resumes stream + callable.call(); + call = innerCallable.popLastCall(); + Truth.assertThat(call.getRequest()).isEqualTo("request > 0"); + } + + @Test + @SuppressWarnings("ConstantConditions") + void testUnauthenticatedRefreshWithNonResumableStreamDoesNotRetry() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + java.util.concurrent.atomic.AtomicLong generation = + new java.util.concurrent.atomic.AtomicLong(0); + Mockito.when(transportChannel.getGeneration()).thenAnswer(inv -> generation.get()); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + Mockito.doAnswer( + inv -> { + generation.incrementAndGet(); + return null; + }) + .when(transportChannel) + .refresh(); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + // SimpleStreamResumptionStrategy cannot resume once a response has been received + resumptionStrategy = new com.google.api.gax.retrying.SimpleStreamResumptionStrategy<>(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + call.getController().getObserver().onResponse("response1"); + + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + ServerStreamingAttemptException attemptEx = (ServerStreamingAttemptException) outerError; + Truth.assertThat(attemptEx.canResume()).isFalse(); + Truth.assertThat(((UnauthenticatedException) attemptEx.getCause()).isChannelRefreshed()) + .isTrue(); + Truth.assertThat( + new com.google.api.gax.retrying.StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm<>(), + new com.google.api.gax.retrying.ExponentialRetryAlgorithm( + RetrySettings.newBuilder().build(), + com.google.api.core.NanoClock.getDefaultClock())) + .shouldRetry(attemptEx, null, fakeRetryingFuture.getAttemptSettings())) + .isFalse(); + } + + @Test + @SuppressWarnings("ConstantConditions") + void testRefreshThrowsException_originalErrorNotLost() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + Mockito.doThrow(new RuntimeException("Refresh error")).when(transportChannel).refresh(); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(outerError.getCause()).isEqualTo(initialError); + } + @Test @SuppressWarnings("ConstantConditions") void testMidRetry() { @@ -403,6 +609,141 @@ public String processResponse(String response) { .containsExactly("first+suffix", "second+suffix", "third+suffix"); } + @Test + void testUnauthenticatedException_whenChannelRefreshes_flagsChannelRefreshed() throws Exception { + FakeChannel fakeChannel = new FakeChannel(); + fakeChannel.setShouldRefresh(true); + ApiCallContext context = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + UnauthenticatedException unauthEx = + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false); + call.getController().getObserver().onError(unauthEx); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); + Throwable cause = ex.getCause().getCause(); + Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) cause).isChannelRefreshed()).isTrue(); + Truth.assertThat(fakeChannel.getRefreshCount()).isEqualTo(1); + } + + @Test + void testUnauthenticatedException_whenChannelRefreshFails_notFlagged() throws Exception { + java.util.concurrent.atomic.AtomicInteger refreshCalls = + new java.util.concurrent.atomic.AtomicInteger(); + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public void refresh() { + refreshCalls.incrementAndGet(); + throw new RuntimeException("Refresh failed"); + } + }; + fakeChannel.setShouldRefresh(true); + ApiCallContext context = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + UnauthenticatedException unauthEx = + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false); + call.getController().getObserver().onError(unauthEx); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Truth.assertThat(refreshCalls.get()).isEqualTo(1); + Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); + Throwable cause = ex.getCause().getCause(); + Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) cause).isRetryable()).isFalse(); + Truth.assertThat(((UnauthenticatedException) cause).isChannelRefreshed()).isFalse(); + } + + @Test + void testUnauthenticated_generationAlreadyAdvanced_flaggedWithoutCheckingShouldRefresh() + throws Exception { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + java.util.concurrent.atomic.AtomicLong generation = + new java.util.concurrent.atomic.AtomicLong(0); + Mockito.when(transportChannel.getGeneration()).thenAnswer(inv -> generation.get()); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + // Another request rotated the channel while this stream was open. + generation.incrementAndGet(); + call.getController() + .getObserver() + .onError( + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false)); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable cause = ex.getCause().getCause(); + Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) cause).isChannelRefreshed()).isTrue(); + Mockito.verify(transportChannel, Mockito.never()).shouldRefresh(); + Mockito.verify(transportChannel, Mockito.never()).refresh(); + } + + @Test + void testUnauthenticated_shouldRefreshThrows_originalErrorPreserved() throws Exception { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + throw new IllegalStateException("Unable to read certificate"); + } + }; + ApiCallContext context = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + UnauthenticatedException unauthEx = + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false); + call.getController().getObserver().onError(unauthEx); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(ex.getCause().getCause()).isSameInstanceAs(unauthEx); + Truth.assertThat(unauthEx.isChannelRefreshed()).isFalse(); + Truth.assertThat(fakeChannel.getRefreshCount()).isEqualTo(0); + } + static class MyStreamResumptionStrategy implements StreamResumptionStrategy { private int responseCount; diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java index 6f9584826893..c46522114318 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java @@ -129,7 +129,7 @@ void testClientStreamingCall() { ClientStreamingCallable callable = stashCallable.withDefaultCallContext(defaultCallContext); callable.clientStreamingCall(observer); - assertSame(observer, stashCallable.getActualObserver()); + org.junit.jupiter.api.Assertions.assertNotNull(stashCallable.getActualObserver()); assertSame(defaultCallContext, stashCallable.getContext()); } @@ -158,7 +158,7 @@ void testClientStreamingCallWithContext() { ClientStreamingCallable callable = stashCallable.withDefaultCallContext(FakeCallContext.createDefault()); callable.clientStreamingCall(observer, context); - assertSame(observer, stashCallable.getActualObserver()); + org.junit.jupiter.api.Assertions.assertNotNull(stashCallable.getActualObserver()); FakeCallContext actualContext = (FakeCallContext) stashCallable.getContext(); assertSame(channel, actualContext.getChannel()); assertSame(credentials, actualContext.getCredentials()); diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/UnauthenticatedExceptionTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/UnauthenticatedExceptionTest.java new file mode 100644 index 000000000000..083de5286277 --- /dev/null +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/UnauthenticatedExceptionTest.java @@ -0,0 +1,183 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.rpc; + +import static com.google.common.truth.Truth.assertThat; + +import com.google.api.gax.rpc.testing.FakeStatusCode; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.ObjectInputStream; +import java.io.ObjectOutputStream; +import java.io.ObjectStreamClass; +import java.io.Serializable; +import java.util.Base64; +import java.util.Collections; +import org.junit.jupiter.api.Test; + +class UnauthenticatedExceptionTest { + + /** + * An {@link UnauthenticatedException} with message "serialized by released gax", no cause, an + * empty stack trace, retryable {@code true} and a {@link SerializableStatusCode} of + * UNAUTHENTICATED, Java-serialized with the released gax 2.87.0 jar. + */ + private static final String SERIALIZED_BY_GAX_2_87_0 = + "rO0ABXNyAC9jb20uZ29vZ2xlLmFwaS5nYXgucnBjLlVuYXV0aGVudGljYXRlZEV4Y2VwdGlvbmC+YDhKxcJlAgAAeHIAI2NvbS5nb29nbGUuYXBpLmdheC5ycGMuQXBpRXhjZXB0aW9uw0h4sCzSVFQCAANaAAlyZXRyeWFibGVMAAxlcnJvckRldGFpbHN0ACVMY29tL2dvb2dsZS9hcGkvZ2F4L3JwYy9FcnJvckRldGFpbHM7TAAKc3RhdHVzQ29kZXQAI0xjb20vZ29vZ2xlL2FwaS9nYXgvcnBjL1N0YXR1c0NvZGU7eHIAGmphdmEubGFuZy5SdW50aW1lRXhjZXB0aW9unl8GRwo0g+UCAAB4cgATamF2YS5sYW5nLkV4Y2VwdGlvbtD9Hz4aOxzEAgAAeHIAE2phdmEubGFuZy5UaHJvd2FibGXVxjUnOXe4ywMABEwABWNhdXNldAAVTGphdmEvbGFuZy9UaHJvd2FibGU7TAANZGV0YWlsTWVzc2FnZXQAEkxqYXZhL2xhbmcvU3RyaW5nO1sACnN0YWNrVHJhY2V0AB5bTGphdmEvbGFuZy9TdGFja1RyYWNlRWxlbWVudDtMABRzdXBwcmVzc2VkRXhjZXB0aW9uc3QAEExqYXZhL3V0aWwvTGlzdDt4cHB0ABpzZXJpYWxpemVkIGJ5IHJlbGVhc2VkIGdheHVyAB5bTGphdmEubGFuZy5TdGFja1RyYWNlRWxlbWVudDsCRio8PP0iOQIAAHhwAAAAAHNyAB9qYXZhLnV0aWwuQ29sbGVjdGlvbnMkRW1wdHlMaXN0ergXtDynnt4CAAB4cHgBcHNyAEpjb20uZ29vZ2xlLmFwaS5nYXgucnBjLlVuYXV0aGVudGljYXRlZEV4Y2VwdGlvblRlc3QkU2VyaWFsaXphYmxlU3RhdHVzQ29kZQAAAAAAAAABAgABTAAEY29kZXQAKExjb20vZ29vZ2xlL2FwaS9nYXgvcnBjL1N0YXR1c0NvZGUkQ29kZTt4cH5yACZjb20uZ29vZ2xlLmFwaS5nYXgucnBjLlN0YXR1c0NvZGUkQ29kZQAAAAAAAAAAEgAAeHIADmphdmEubGFuZy5FbnVtAAAAAAAAAAASAAB4cHQAD1VOQVVUSEVOVElDQVRFRA=="; + + /** A serializable {@link StatusCode}, since the transport implementations are not. */ + static final class SerializableStatusCode implements StatusCode, Serializable { + private static final long serialVersionUID = 1L; + private final Code code; + + SerializableStatusCode(Code code) { + this.code = code; + } + + @Override + public Code getCode() { + return code; + } + + @Override + public Object getTransportCode() { + return code.name(); + } + } + + @Test + void publicConstructors_areNotChannelRefreshed() { + StatusCode statusCode = FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED); + ErrorDetails errorDetails = + ErrorDetails.builder().setRawErrorMessages(Collections.emptyList()).build(); + + assertThat(new UnauthenticatedException(null, statusCode, true).isChannelRefreshed()).isFalse(); + assertThat(new UnauthenticatedException("msg", null, statusCode, true).isChannelRefreshed()) + .isFalse(); + assertThat( + new UnauthenticatedException(null, statusCode, true, errorDetails).isChannelRefreshed()) + .isFalse(); + assertThat( + new UnauthenticatedException("msg", null, statusCode, true, errorDetails) + .isChannelRefreshed()) + .isFalse(); + } + + @Test + void withChannelRefreshed_preservesFields() { + ErrorDetails errorDetails = + ErrorDetails.builder().setRawErrorMessages(Collections.emptyList()).build(); + IllegalStateException cause = new IllegalStateException("root cause"); + StatusCode statusCode = FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED); + UnauthenticatedException original = + new UnauthenticatedException("Expired cert", cause, statusCode, false, errorDetails); + original.setStackTrace( + new StackTraceElement[] {new StackTraceElement("foo", "bar", "Baz.java", 123)}); + RuntimeException suppressed = new RuntimeException("suppressed"); + original.addSuppressed(suppressed); + + UnauthenticatedException refreshed = original.withChannelRefreshed(); + + assertThat(refreshed).isNotSameInstanceAs(original); + assertThat(refreshed.isChannelRefreshed()).isTrue(); + assertThat(original.isChannelRefreshed()).isFalse(); + assertThat(refreshed.getMessage()).isEqualTo(original.getMessage()); + assertThat(refreshed.getCause()).isSameInstanceAs(cause); + assertThat(refreshed.getStatusCode()).isSameInstanceAs(statusCode); + assertThat(refreshed.getErrorDetails()).isSameInstanceAs(errorDetails); + assertThat(refreshed.getStackTrace()).isEqualTo(original.getStackTrace()); + assertThat(refreshed.getSuppressed()).asList().containsExactly(suppressed); + } + + @Test + void withChannelRefreshed_keepsIsRetryable() { + StatusCode statusCode = FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED); + + assertThat( + new UnauthenticatedException("msg", null, statusCode, false) + .withChannelRefreshed() + .isRetryable()) + .isFalse(); + assertThat( + new UnauthenticatedException("msg", null, statusCode, true) + .withChannelRefreshed() + .isRetryable()) + .isTrue(); + } + + @Test + void serialVersionUID_matchesReleasedValue() { + // gax 2.83.0 to 2.87.0 did not declare a serialVersionUID; this is the value the JVM computed + // for them. Keeping it avoids InvalidClassException across versions. + assertThat(ObjectStreamClass.lookup(UnauthenticatedException.class).getSerialVersionUID()) + .isEqualTo(6971115068105015909L); + } + + @Test + void deserializesExceptionSerializedByReleasedGax() throws Exception { + Object deserialized = deserialize(Base64.getDecoder().decode(SERIALIZED_BY_GAX_2_87_0)); + + assertThat(deserialized).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException ex = (UnauthenticatedException) deserialized; + assertThat(ex.getMessage()).isEqualTo("serialized by released gax"); + assertThat(ex.getStatusCode().getCode()).isEqualTo(StatusCode.Code.UNAUTHENTICATED); + assertThat(ex.isRetryable()).isTrue(); + assertThat(ex.isChannelRefreshed()).isFalse(); + } + + @Test + void withChannelRefreshed_flagNotSerialized() throws Exception { + UnauthenticatedException refreshed = + new UnauthenticatedException( + "msg", null, new SerializableStatusCode(StatusCode.Code.UNAUTHENTICATED), false) + .withChannelRefreshed(); + + UnauthenticatedException roundTripped = + (UnauthenticatedException) deserialize(serialize(refreshed)); + + assertThat(roundTripped.getMessage()).isEqualTo("msg"); + assertThat(roundTripped.isRetryable()).isFalse(); + assertThat(roundTripped.isChannelRefreshed()).isFalse(); + } + + private static byte[] serialize(Object object) throws Exception { + ByteArrayOutputStream bytes = new ByteArrayOutputStream(); + try (ObjectOutputStream out = new ObjectOutputStream(bytes)) { + out.writeObject(object); + } + return bytes.toByteArray(); + } + + private static Object deserialize(byte[] bytes) throws Exception { + try (ObjectInputStream in = new ObjectInputStream(new ByteArrayInputStream(bytes))) { + return in.readObject(); + } + } +} diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java index bea4674b765b..5ed563254965 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java @@ -64,8 +64,10 @@ void testNotUseClientCertificate() throws IOException, GeneralSecurityException @Test void testUseClientCertificate() throws IOException, GeneralSecurityException { CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); MtlsProvider provider = new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); assertNotNull(getMtlsObjectFromTransportChannel(provider, certificateBasedAccess)); @@ -74,8 +76,10 @@ void testUseClientCertificate() throws IOException, GeneralSecurityException { @Test void testNoClientCertificate() throws IOException, GeneralSecurityException { CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); MtlsProvider provider = new FakeMtlsProvider(null, "", false); assertNull(getMtlsObjectFromTransportChannel(provider, certificateBasedAccess)); } @@ -84,8 +88,10 @@ void testNoClientCertificate() throws IOException, GeneralSecurityException { void testGetKeyStoreThrows() throws GeneralSecurityException { // Test the case where provider.getKeyStore() throws. CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); MtlsProvider provider = new FakeMtlsProvider(null, "", true); IOException actual = assertThrows( diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index e328e0af4799..92c0dfb96f3a 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -32,52 +32,131 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +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 java.util.HashMap; +import java.util.Map; import org.junit.jupiter.api.Test; class CertificateBasedAccessTest { + private static class TestEnv { + private final Map env = new HashMap<>(); + + TestEnv() { + // Hermetically isolate tests from the host's ~/.config/gcloud/certificate_config.json + env.put("CLOUDSDK_CONFIG", "/nonexistent/test/gcloud"); + } + + void set(String key, String val) { + env.put(key, val); + } + + String get(String name) { + return env.get(name); + } + } + + private CertificateBasedAccess createCba(TestEnv env) { + return new CertificateBasedAccess(env::get); + } + @Test void testUseMtlsEndpointAlways() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "false"); + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "always"); + CertificateBasedAccess cba = createCba(env); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS, cba.getMtlsEndpointUsagePolicy()); } @Test void testUseMtlsEndpointAuto() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "false"); + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "auto"); + CertificateBasedAccess cba = createCba(env); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO, cba.getMtlsEndpointUsagePolicy()); } @Test void testUseMtlsEndpointNever() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "never" : "false"); + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "never"); + CertificateBasedAccess cba = createCba(env); + assertEquals( + CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER, cba.getMtlsEndpointUsagePolicy()); + } + + @Test + void testUseMtlsEndpointCaseInsensitive() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "ALWAYS"); + CertificateBasedAccess cba = createCba(env); + assertEquals( + CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS, cba.getMtlsEndpointUsagePolicy()); + + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "NEVER"); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER, cba.getMtlsEndpointUsagePolicy()); } @Test - void testUseMtlsClientCertificateTrue() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_CLIENT_CERTIFICATE") ? "true" : "auto"); + void testUseMtlsClientCertificateExplicitTrueNoCredentials() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); + CertificateBasedAccess cba = createCba(env); + // Explicit 'true' enables mTLS client certificate usage (for ECP / custom MtlsProvider) even + // when no workload cert files are present, while getWorkloadCertPath returns null so file + // rotation polling is not active. assertTrue(cba.useMtlsClientCertificate()); + assertNull(cba.getWorkloadCertPath()); } @Test - void testUseMtlsClientCertificateFalse() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_CLIENT_CERTIFICATE") ? "false" : "auto"); + void testUseMtlsClientCertificateExplicitFalse() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "false"); + + CertificateBasedAccess cba = createCba(env); + assertFalse(cba.useMtlsClientCertificate()); + assertNull(cba.getWorkloadCertPath()); + } + + @Test + void testUseMtlsClientCertificateUnsetNoFiles() { + TestEnv env = new TestEnv(); + CertificateBasedAccess cba = createCba(env); assertFalse(cba.useMtlsClientCertificate()); + assertNull(cba.getWorkloadCertPath()); + } + + @Test + void testUseMtlsClientCertificateConfigMissingConfigFile_throwsIllegalStateException() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/nonexistent/config.json"); + + CertificateBasedAccess cba = createCba(env); + + // Non-existent config file on disk specified via explicit env var throws IllegalStateException + // (Fail Closed) + assertThrows(IllegalStateException.class, () -> cba.useMtlsClientCertificate()); + assertThrows(IllegalStateException.class, () -> cba.getWorkloadCertPath()); + } + + @Test + void testWorkloadCertificateUtilsEmptyFileReturnsEmptyString() throws Exception { + java.io.File tempFile = java.io.File.createTempFile("test-cert-empty", ".pem"); + tempFile.deleteOnExit(); + // 0-byte truncated file mid-write should return empty string rather than SHA-256 of empty bytes + assertEquals( + "", WorkloadCertificateUtils.getCertificateFingerprint(tempFile.getAbsolutePath())); + + java.nio.file.Files.write( + tempFile.toPath(), "test-cert-content".getBytes(java.nio.charset.StandardCharsets.UTF_8)); + String fp = WorkloadCertificateUtils.getCertificateFingerprint(tempFile.getAbsolutePath()); + assertFalse(fp.isEmpty()); } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java index 1cdefe435d55..e41dc041220d 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java @@ -241,6 +241,11 @@ public FakeChannel getChannel() { return channel; } + @Override + public TransportChannel getTransportChannel() { + return channel != null ? FakeTransportChannel.create(channel) : null; + } + @Override public java.time.Duration getTimeoutDuration() { return timeout; diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java index b725da363b3b..9acfc703b9c3 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java @@ -32,4 +32,36 @@ import com.google.api.core.InternalApi; @InternalApi("for testing") -public class FakeChannel {} +public class FakeChannel { + private volatile boolean shouldRefresh = false; + private volatile int refreshCount = 0; + + public FakeChannel setShouldRefresh(boolean shouldRefresh) { + this.shouldRefresh = shouldRefresh; + return this; + } + + public boolean shouldRefresh() { + return shouldRefresh; + } + + public void refresh() { + refreshCount++; + generation++; + } + + public int getRefreshCount() { + return refreshCount; + } + + private volatile long generation = 0; + + public FakeChannel setGeneration(long generation) { + this.generation = generation; + return this; + } + + public long getGeneration() { + return generation; + } +} diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java index 0d4abac8f1c6..7d92556b88d5 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java @@ -41,6 +41,48 @@ public class FakeTransportChannel implements TransportChannel { private volatile boolean isShutdown = false; private volatile Map headers; private volatile Executor executor; + private volatile boolean shouldRefresh = false; + private volatile int refreshCount = 0; + + public FakeTransportChannel setShouldRefresh(boolean shouldRefresh) { + if (channel != null) { + channel.setShouldRefresh(shouldRefresh); + } + this.shouldRefresh = shouldRefresh; + return this; + } + + @Override + public boolean shouldRefresh() { + return channel != null ? channel.shouldRefresh() : shouldRefresh; + } + + @Override + public void refresh() { + if (channel != null) { + channel.refresh(); + } + refreshCount++; + } + + public int getRefreshCount() { + return channel != null ? channel.getRefreshCount() : refreshCount; + } + + private volatile long generation = 0; + + public FakeTransportChannel setGeneration(long generation) { + if (channel != null) { + channel.setGeneration(generation); + } + this.generation = generation; + return this; + } + + @Override + public long getGeneration() { + return channel != null ? channel.getGeneration() : generation; + } private FakeTransportChannel(FakeChannel channel) { this.channel = channel;