Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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<String>): Array<String> {
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 {
Expand All @@ -66,4 +86,61 @@ object HttpClientFactory {
}
}
}

private class ProtocolFilteringSslSocketFactory(
private val delegate: SSLSocketFactory
) : SSLSocketFactory() {

override fun getDefaultCipherSuites(): Array<String> {
return delegate.defaultCipherSuites
}

override fun getSupportedCipherSuites(): Array<String> {
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
}
}
}
}
Original file line number Diff line number Diff line change
@@ -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"
)
)
}
}
Loading