diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactory.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactory.java index 136ce607..6ade938a 100644 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactory.java +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactory.java @@ -23,6 +23,7 @@ import org.apache.commons.logging.Log; /** * Default implementation of {@link ApacheHttpClientConnectionManagerFactory}. * @author Ryan Baxter + * @author Michael Wirth */ public class DefaultApacheHttpClientConnectionManagerFactory implements ApacheHttpClientConnectionManagerFactory { @@ -47,24 +48,11 @@ public class DefaultApacheHttpClientConnectionManagerFactory if (disableSslValidation) { try { final SSLContext sslContext = SSLContext.getInstance("SSL"); - sslContext.init(null, new TrustManager[] { new X509TrustManager() { - @Override - public void checkClientTrusted(X509Certificate[] x509Certificates, - String s) throws CertificateException { - } - - @Override - public void checkServerTrusted(X509Certificate[] x509Certificates, - String s) throws CertificateException { - } - - @Override - public X509Certificate[] getAcceptedIssuers() { - return null; - } - } }, new SecureRandom()); + sslContext.init(null, + new TrustManager[] { new DisabledValidationTrustManager()}, + new SecureRandom()); registryBuilder.register(HTTPS_SCHEME, new SSLConnectionSocketFactory( - sslContext, NoopHostnameVerifier.INSTANCE)); + sslContext, NoopHostnameVerifier.INSTANCE)); } catch (NoSuchAlgorithmException e) { LOG.warn("Error creating SSLContext", e); @@ -72,6 +60,8 @@ public class DefaultApacheHttpClientConnectionManagerFactory catch (KeyManagementException e) { LOG.warn("Error creating SSLContext", e); } + } else { + registryBuilder.register("https", SSLConnectionSocketFactory.getSocketFactory()); } final Registry registry = registryBuilder.build(); @@ -82,4 +72,21 @@ public class DefaultApacheHttpClientConnectionManagerFactory return connectionManager; } + + class DisabledValidationTrustManager implements X509TrustManager { + @Override + public void checkClientTrusted(X509Certificate[] x509Certificates, + String s) throws CertificateException { + } + + @Override + public void checkServerTrusted(X509Certificate[] x509Certificates, + String s) throws CertificateException { + } + + @Override + public X509Certificate[] getAcceptedIssuers() { + return null; + } + } } diff --git a/spring-cloud-commons/src/test/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactoryTests.java b/spring-cloud-commons/src/test/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactoryTests.java index 2544f6ca..b5433fac 100644 --- a/spring-cloud-commons/src/test/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactoryTests.java +++ b/spring-cloud-commons/src/test/java/org/springframework/cloud/commons/httpclient/DefaultApacheHttpClientConnectionManagerFactoryTests.java @@ -1,17 +1,26 @@ package org.springframework.cloud.commons.httpclient; -import java.lang.reflect.Field; -import java.util.concurrent.TimeUnit; - +import org.apache.http.config.Lookup; import org.apache.http.conn.HttpClientConnectionManager; +import org.apache.http.conn.socket.ConnectionSocketFactory; +import org.apache.http.impl.conn.DefaultHttpClientConnectionOperator; import org.apache.http.impl.conn.PoolingHttpClientConnectionManager; import org.junit.Test; import org.springframework.util.ReflectionUtils; +import javax.net.ssl.SSLContextSpi; +import javax.net.ssl.SSLSocketFactory; +import javax.net.ssl.X509TrustManager; +import java.lang.reflect.Field; +import java.util.concurrent.TimeUnit; + +import static org.hamcrest.Matchers.*; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertThat; /** * @author Ryan Baxter + * @author Michael Wirth */ public class DefaultApacheHttpClientConnectionManagerFactoryTests { @Test @@ -42,6 +51,42 @@ public class DefaultApacheHttpClientConnectionManagerFactoryTests { assertEquals(TimeUnit.DAYS, timeUnit); } + @Test + public void newConnectionManagerWithSSL() throws Exception { + HttpClientConnectionManager connectionManager = new DefaultApacheHttpClientConnectionManagerFactory() + .newConnectionManager(false, 2, 6); + + Lookup socketFactoryRegistry = getConnectionSocketFactoryLookup( + connectionManager); + assertThat(socketFactoryRegistry.lookup("https"), is(notNullValue())); + assertThat(getX509TrustManager(socketFactoryRegistry).getAcceptedIssuers(), is(notNullValue())); + } + + @Test + public void newConnectionManagerWithDisabledSSLValidation() throws Exception { + HttpClientConnectionManager connectionManager = new DefaultApacheHttpClientConnectionManagerFactory() + .newConnectionManager(true, 2, 6); + + Lookup socketFactoryRegistry = getConnectionSocketFactoryLookup( + connectionManager); + assertThat(socketFactoryRegistry.lookup("https"), is(notNullValue())); + assertThat(getX509TrustManager(socketFactoryRegistry).getAcceptedIssuers(), is(nullValue())); + } + + private Lookup getConnectionSocketFactoryLookup( + HttpClientConnectionManager connectionManager) { + DefaultHttpClientConnectionOperator connectionOperator = getField(connectionManager, "connectionOperator"); + return getField(connectionOperator, "socketFactoryRegistry"); + } + + private X509TrustManager getX509TrustManager( + Lookup socketFactoryRegistry) { + ConnectionSocketFactory connectionSocketFactory = socketFactoryRegistry.lookup("https"); + SSLSocketFactory sslSocketFactory = getField(connectionSocketFactory, "socketfactory"); + SSLContextSpi sslContext = getField(sslSocketFactory, "context"); + return getField(sslContext, "trustManager"); + } + @SuppressWarnings("unchecked") protected T getField(Object target, String name) { Field field = ReflectionUtils.findField(target.getClass(), name);