diff --git a/app/src/main/kotlin/com/ivanovsky/passnotes/data/repository/file/webdav/HttpClientFactory.kt b/app/src/main/kotlin/com/ivanovsky/passnotes/data/repository/file/webdav/HttpClientFactory.kt index fdc38021..f245a22d 100644 --- a/app/src/main/kotlin/com/ivanovsky/passnotes/data/repository/file/webdav/HttpClientFactory.kt +++ b/app/src/main/kotlin/com/ivanovsky/passnotes/data/repository/file/webdav/HttpClientFactory.kt @@ -2,9 +2,13 @@ package com.ivanovsky.passnotes.data.repository.file.webdav import android.annotation.SuppressLint import com.ivanovsky.passnotes.BuildConfig +import java.net.InetAddress +import java.net.Socket import java.security.SecureRandom import java.security.cert.X509Certificate import javax.net.ssl.SSLContext +import javax.net.ssl.SSLSocket +import javax.net.ssl.SSLSocketFactory import javax.net.ssl.X509TrustManager import okhttp3.OkHttp import okhttp3.OkHttpClient @@ -37,13 +41,29 @@ object HttpClientFactory { val sslContext = SSLContext.getInstance("TLS") sslContext.init(null, arrayOf(unsecuredTrustManager), SecureRandom()) - builder.sslSocketFactory(sslContext.socketFactory, unsecuredTrustManager) + val protocolFilteringSocketFactory = ProtocolFilteringSslSocketFactory(sslContext.socketFactory) + builder.sslSocketFactory(protocolFilteringSocketFactory, unsecuredTrustManager) builder.hostnameVerifier { _, _ -> true } } return builder.build() } + internal fun filterEnabledProtocols(enabledProtocols: Array): Array { + val filteredProtocols = enabledProtocols.filterNot { protocol -> + protocol.startsWith("SSL", ignoreCase = true) || + protocol.equals("TLSv1", ignoreCase = true) || + protocol.equals("TLSv1.0", ignoreCase = true) || + protocol.equals("TLSv1.1", ignoreCase = true) + } + + check(filteredProtocols.isNotEmpty()) { + "No acceptable TLS protocols are enabled by the SSL socket provider" + } + + return filteredProtocols.toTypedArray() + } + @SuppressLint("CustomX509TrustManager") private fun createUnsecuredTrustManager(): X509TrustManager { return object : X509TrustManager { @@ -66,4 +86,61 @@ object HttpClientFactory { } } } + + private class ProtocolFilteringSslSocketFactory( + private val delegate: SSLSocketFactory + ) : SSLSocketFactory() { + + override fun getDefaultCipherSuites(): Array { + return delegate.defaultCipherSuites + } + + override fun getSupportedCipherSuites(): Array { + return delegate.supportedCipherSuites + } + + override fun createSocket(): Socket { + return configure(delegate.createSocket()) + } + + override fun createSocket(host: String, port: Int): Socket { + return configure(delegate.createSocket(host, port)) + } + + override fun createSocket(host: String, port: Int, localHost: InetAddress, localPort: Int): Socket { + return configure(delegate.createSocket(host, port, localHost, localPort)) + } + + override fun createSocket(host: InetAddress, port: Int): Socket { + return configure(delegate.createSocket(host, port)) + } + + override fun createSocket( + address: InetAddress, + port: Int, + localAddress: InetAddress, + localPort: Int + ): Socket { + return configure(delegate.createSocket(address, port, localAddress, localPort)) + } + + override fun createSocket(socket: Socket, host: String, port: Int, autoClose: Boolean): Socket { + return configure(delegate.createSocket(socket, host, port, autoClose)) + } + + private fun configure(socket: Socket): Socket { + try { + val sslSocket = socket as? SSLSocket + ?: throw IllegalStateException("SSL socket factory returned a non-SSL socket") + sslSocket.enabledProtocols = filterEnabledProtocols(sslSocket.enabledProtocols) + return sslSocket + } catch (exception: Exception) { + try { + socket.close() + } catch (_: Exception) { + } + throw exception + } + } + } } \ No newline at end of file diff --git a/app/src/test/kotlin/com/ivanovsky/passnotes/data/repository/file/webdav/HttpClientFactoryTest.kt b/app/src/test/kotlin/com/ivanovsky/passnotes/data/repository/file/webdav/HttpClientFactoryTest.kt new file mode 100644 index 00000000..1919ab9d --- /dev/null +++ b/app/src/test/kotlin/com/ivanovsky/passnotes/data/repository/file/webdav/HttpClientFactoryTest.kt @@ -0,0 +1,50 @@ +package com.ivanovsky.passnotes.data.repository.file.webdav + +import org.junit.Assert.assertEquals +import org.junit.Test + +class HttpClientFactoryTest { + + @Test + fun `filterEnabledProtocols removes SSL and legacy TLS protocols while retaining TLSv1 2 and TLSv1 3`() { + val filteredProtocols = HttpClientFactory.filterEnabledProtocols( + arrayOf( + "SSL", + "SSLv2", + "SSLv2Hello", + "SSLv3", + "TLSv1", + "TLSv1.0", + "TLSv1.1", + "TLSv1.2", + "TLSv1.3" + ) + ) + + assertEquals(listOf("TLSv1.2", "TLSv1.3"), filteredProtocols.toList()) + } + + @Test + fun `filterEnabledProtocols retains an unknown future TLS protocol`() { + val filteredProtocols = HttpClientFactory.filterEnabledProtocols( + arrayOf("TLSv1.4") + ) + + assertEquals(listOf("TLSv1.4"), filteredProtocols.toList()) + } + + @Test(expected = IllegalStateException::class) + fun `filterEnabledProtocols fails closed when only legacy protocols are enabled`() { + HttpClientFactory.filterEnabledProtocols( + arrayOf( + "SSL", + "SSLv2", + "SSLv2Hello", + "SSLv3", + "TLSv1", + "TLSv1.0", + "TLSv1.1" + ) + ) + } +}