Skip to content
Merged
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 @@ -17,12 +17,9 @@ package okhttp.android.test

import android.annotation.SuppressLint
import android.net.Network
import android.os.Build
import java.net.InetAddress
import okhttp3.Dns
import okhttp3.Interceptor
import okhttp3.Response
import okhttp3.android.EchAwareDns
import okhttp3.android.AndroidDns

/**
* Interceptor that supports Network Pinning on Android via Request tags.
Expand All @@ -43,34 +40,11 @@ class AndroidNetworkPinning : Interceptor {
if (pinnedNetwork != null) {
chain
.withSocketFactory(pinnedNetwork.socketFactory)
.withDns(dnsForNetwork(pinnedNetwork))
.withDns(AndroidDns(network = pinnedNetwork))
} else {
chain
}

return effectiveChain.proceed(request)
}

/**
* ECH needs the `HTTPS` record, which is only reachable through `DnsResolver.rawQuery()` and only
* consulted by the platform from API 37. Below that there's nothing to gain from the extra query,
* so [AndroidNetworkDns] does the plain address lookup.
*/
private fun dnsForNetwork(network: Network): Dns =
when {
Build.VERSION.SDK_INT >= 37 -> EchAwareDns.forNetwork(network)
else -> AndroidNetworkDns(network)
}
}

/**
* A [Dns] scoped to [network], used below API 37 where there's no ECH to resolve for.
*
* [Network.getAllByName] is the whole implementation: it resolves on that network and nothing else,
* with no service metadata.
*/
class AndroidNetworkDns(
private val network: Network,
) : Dns {
override fun lookup(hostname: String): List<InetAddress> = network.getAllByName(hostname).toList()
}
67 changes: 29 additions & 38 deletions android-test/src/androidTest/java/okhttp/android/test/EchTest.kt
Original file line number Diff line number Diff line change
Expand Up @@ -25,10 +25,11 @@ import assertk.assertions.doesNotContain
import assertk.assertions.isEqualTo
import assertk.assertions.isFalse
import assertk.assertions.isTrue
import okhttp3.Dns
import okhttp3.HttpUrl.Companion.toHttpUrl
import okhttp3.OkHttpClient
import okhttp3.Request
import okhttp3.android.EchAwareDns
import okhttp3.android.AndroidDns
import okhttp3.dnsoverhttps.DnsOverHttps
import org.junit.jupiter.api.Assumptions.assumeTrue
import org.junit.jupiter.api.BeforeEach
Expand All @@ -37,29 +38,30 @@ import org.junit.jupiter.api.Test
import org.junit.jupiter.api.fail

/**
* Confirms Encrypted Client Hello (ECH) end to end, with [EchAwareDns].
* Confirms Encrypted Client Hello (ECH) end to end.
*
* Test with both [okhttp3.android.AndroidDns] and [DnsOverHttps].
* Test with both [AndroidDns] and [DnsOverHttps].
*
* See `res/xml/network_security_config.xml` for overrides.
*/
@SuppressLint("NewApi")
@Tag("Remote")
@Burst
class EchTest(
private val useDoh: Boolean = false,
private val dnsApi: DnsApi = DnsApi.Doh,
) {
private lateinit var client: OkHttpClient

@BeforeEach
fun setUp() {
// EchAwareDns reads API 37 NetworkSecurityPolicy.getDomainEncryptionMode().
// ECH requires API 37.
assumeTrue(Build.VERSION.SDK_INT >= 37)

val bootstrapClient = OkHttpClient()
val dns = dnsApi.create(bootstrapClient)
client =
OkHttpClient
.Builder()
.dns(dns())
bootstrapClient.newBuilder()
.dns(dns)
.build()
}

Expand Down Expand Up @@ -131,36 +133,25 @@ class EchTest(
assertThat(client.get("https://crypto.cloudflare.com/cdn-cgi/trace")).contains("sni=plaintext")
}

/**
* [EchAwareDns] over the platform resolver, or over DoH when [useDoh]. Both arms use the same
* source: the ECH one carries service metadata, the other doesn't.
*/
private fun dns(): EchAwareDns =
when {
useDoh -> {
val bootstrapClient = OkHttpClient()
EchAwareDns(
echDns = dnsOverHttps(bootstrapClient, includeServiceMetadata = true),
addressOnlyDns = dnsOverHttps(bootstrapClient, includeServiceMetadata = false),
)
}
else -> EchAwareDns()
}

/** Addressed by IP, so resolving the resolver doesn't need a resolver. */
private fun dnsOverHttps(
bootstrapClient: OkHttpClient,
includeServiceMetadata: Boolean,
): DnsOverHttps =
DnsOverHttps
.Builder()
.client(bootstrapClient)
.url("https://1.1.1.1/dns-query".toHttpUrl())
.includeServiceMetadata(includeServiceMetadata)
.build()

private fun OkHttpClient.get(url: String): String =
newCall(Request.Builder().url(url).build()).execute().use { response ->
fun OkHttpClient.get(url: String): String =
newCall(Request(url.toHttpUrl())).execute().use { response ->
response.body.string()
}

enum class DnsApi {
Android {
override fun create(client: OkHttpClient) = AndroidDns()
},

Doh {
/** DNS server is addressed by IP, so resolving the resolver doesn't need a resolver. */
override fun create(client: OkHttpClient) = DnsOverHttps.Builder()
.client(client)
.url("https://1.1.1.1/dns-query".toHttpUrl())
.includeServiceMetadata(true)
.build()
};

abstract fun create(client: OkHttpClient): Dns
}
}
2 changes: 1 addition & 1 deletion android-test/src/main/res/xml/network_security_config.xml
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
<domain-config cleartextTrafficPermitted="true">
<domain includeSubdomains="false">localhost</domain>
</domain-config>
<!-- ECH test servers. EchAwareDns queries the HTTPS record. -->
<!-- ECH test servers. AndroidDns queries the HTTPS record. -->
<domain-config>
<domain includeSubdomains="true">cloudflare-ech.com</domain>
<domain includeSubdomains="true">tls-ech.dev</domain>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,7 @@ import okhttp3.Dns
import okhttp3.DnsCache
import okhttp3.EventRecorder
import okhttp3.FakeDns
import okhttp3.FakeDns.Request.DnsOverHttpsRequest
import okhttp3.FakeDns.Request.DnsRequest
import okhttp3.Headers.Companion.headersOf
import okhttp3.Interceptor
import okhttp3.OkHttpClient
Expand Down Expand Up @@ -151,8 +151,8 @@ class DnsOverHttpsTest(
server["lysine.dev"] = listOf(InetAddress.getByName("10.20.30.40"))
val result = dns.invoke(entryPoint, "lysine.dev")
assertThat(result).isEqualTo(listOf(address("10.20.30.40")))
val (httpsRequest, dnsRequest) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest.method).isEqualTo("GET")
val (dnsRequest, httpsRequest) = server.takeRequest() as DnsRequest
assertThat(httpsRequest!!.method).isEqualTo("GET")
assertThat(dnsRequest)
.isEqualTo(queryRequest("lysine.dev", TYPE_A))
}
Expand All @@ -166,8 +166,8 @@ class DnsOverHttpsTest(
server["lysine.dev"] = listOf(InetAddress.getByName("10.20.30.40"))
val result0 = dns.invoke(entryPoint, "lysine.dev")
assertThat(result0).isEqualTo(listOf(address("10.20.30.40")))
val (httpsRequest, dnsRequest) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest.method).isEqualTo("GET")
val (dnsRequest, httpsRequest) = server.takeRequest() as DnsRequest
assertThat(httpsRequest!!.method).isEqualTo("GET")
assertThat(dnsRequest)
.isEqualTo(queryRequest("lysine.dev", TYPE_A))

Expand All @@ -193,12 +193,12 @@ class DnsOverHttpsTest(
address("10.20.30.40"),
)

val (httpsRequest1, dnsRequest1) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest1.method).isEqualTo("GET")
val (dnsRequest1, httpsRequest1) = server.takeRequest() as DnsRequest
assertThat(httpsRequest1!!.method).isEqualTo("GET")
assertThat(dnsRequest1).isEqualTo(queryRequest("lysine.dev", TYPE_AAAA))

val (httpsRequest2, dnsRequest2) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest2.method).isEqualTo("GET")
val (dnsRequest2, httpsRequest2) = server.takeRequest() as DnsRequest
assertThat(httpsRequest2!!.method).isEqualTo("GET")
assertThat(dnsRequest2).isEqualTo(queryRequest("lysine.dev", TYPE_A))
}

Expand All @@ -218,10 +218,10 @@ class DnsOverHttpsTest(
address("10.20.30.40"),
)

val (_, dnsRequest1) = server.takeRequest() as DnsOverHttpsRequest
val (dnsRequest1, _) = server.takeRequest() as DnsRequest
assertThat(dnsRequest1).isEqualTo(queryRequest("lysine.dev", TYPE_AAAA))

val (_, dnsRequest2) = server.takeRequest() as DnsOverHttpsRequest
val (dnsRequest2, _) = server.takeRequest() as DnsRequest
assertThat(dnsRequest2).isEqualTo(queryRequest("lysine.dev", TYPE_A))

assertThat(server.pollRequest()).isNull()
Expand All @@ -232,8 +232,8 @@ class DnsOverHttpsTest(
assertFailsWith<UnknownHostException> {
dns(entryPoint, "lysine.dev")
}
val (httpsRequest, dnsRequest) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest.method).isEqualTo("GET")
val (dnsRequest, httpsRequest) = server.takeRequest() as DnsRequest
assertThat(httpsRequest!!.method).isEqualTo("GET")
assertThat(dnsRequest)
.isEqualTo(queryRequest("lysine.dev", TYPE_A))
}
Expand Down Expand Up @@ -335,8 +335,8 @@ class DnsOverHttpsTest(

val result1 = cachedDns(entryPoint, "lysine.dev")
assertThat(result1).containsExactly(address("10.20.30.40"))
val (httpsRequest1, dnsRequest1) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest1.method).isEqualTo("GET")
val (dnsRequest1, httpsRequest1) = server.takeRequest() as DnsRequest
assertThat(httpsRequest1!!.method).isEqualTo("GET")
assertThat(dnsRequest1)
.isEqualTo(queryRequest("lysine.dev", TYPE_A))

Expand All @@ -350,8 +350,8 @@ class DnsOverHttpsTest(

val result3 = cachedDns(entryPoint, "alternate.lysine.dev")
assertThat(result3).containsExactly(address("55.66.77.88"))
val (httpsRequest2, dnsRequest2) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest2.method).isEqualTo("GET")
val (dnsRequest2, httpsRequest2) = server.takeRequest() as DnsRequest
assertThat(httpsRequest2!!.method).isEqualTo("GET")
assertThat(dnsRequest2)
.isEqualTo(queryRequest("alternate.lysine.dev", TYPE_A))

Expand All @@ -378,8 +378,8 @@ class DnsOverHttpsTest(

val result1 = cachedDns(entryPoint, "lysine.dev")
assertThat(result1).containsExactly(address("10.20.30.40"))
val (httpsRequest1, _) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest1.method).isEqualTo("POST")
val (_, httpsRequest1) = server.takeRequest() as DnsRequest
assertThat(httpsRequest1!!.method).isEqualTo("POST")
assertThat(httpsRequest1.url.encodedQuery)
.isEqualTo("ct")

Expand All @@ -393,8 +393,8 @@ class DnsOverHttpsTest(

val result3 = cachedDns(entryPoint, "alternate.lysine.dev")
assertThat(result3).containsExactly(address("55.66.77.88"))
val (httpsRequest2, _) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest2.method).isEqualTo("POST")
val (_, httpsRequest2) = server.takeRequest() as DnsRequest
assertThat(httpsRequest2!!.method).isEqualTo("POST")
assertThat(httpsRequest2.url.encodedQuery)
.isEqualTo("ct")

Expand All @@ -418,17 +418,17 @@ class DnsOverHttpsTest(
server["lysine.dev"] = listOf(InetAddress.getByName("10.20.30.40"))
val result1 = cachedDns(entryPoint, "lysine.dev")
assertThat(result1).containsExactly(address("10.20.30.40"))
val (httpsRequest1, dnsRequest1) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest1.method).isEqualTo("GET")
val (dnsRequest1, httpsRequest1) = server.takeRequest() as DnsRequest
assertThat(httpsRequest1!!.method).isEqualTo("GET")
assertThat(dnsRequest1)
.isEqualTo(queryRequest("lysine.dev", TYPE_A))

assertThat(cacheEvents()).containsExactly(CacheMiss::class)

val result2 = cachedDns(entryPoint, "lysine.dev")
assertThat(result2).isEqualTo(listOf(address("10.20.30.40")))
val (httpsRequest2, dnsRequest2) = server.takeRequest() as DnsOverHttpsRequest
assertThat(httpsRequest2.method).isEqualTo("GET")
val (dnsRequest2, httpsRequest2) = server.takeRequest() as DnsRequest
assertThat(httpsRequest2!!.method).isEqualTo("GET")
assertThat(dnsRequest2)
.isEqualTo(queryRequest("lysine.dev", TYPE_A))

Expand Down
13 changes: 9 additions & 4 deletions okhttp-testing-support/src/main/kotlin/okhttp3/FakeDns.kt
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ class FakeDns(
}

val dnsRequest = DnsMessageReader(encodedDnsQuery).read()
requests.put(Request.DnsOverHttpsRequest(request, dnsRequest))
requests.put(Request.DnsRequest(dnsRequest, request))

val dnsResponse = invoke(dnsRequest)

Expand Down Expand Up @@ -183,7 +183,12 @@ class FakeDns(
}
}

fun invoke(request: DnsMessage): DnsMessage {
fun query(request: DnsMessage): DnsMessage {
requests.put(Request.DnsRequest(request))
return invoke(request)
}

private fun invoke(request: DnsMessage): DnsMessage {
val answers =
buildList {
for (question in request.questions) {
Expand Down Expand Up @@ -314,9 +319,9 @@ class FakeDns(
sealed interface Request {
val hostname: String

data class DnsOverHttpsRequest(
val httpRequest: RecordedRequest,
data class DnsRequest(
val dnsRequest: DnsMessage,
val httpRequest: RecordedRequest? = null,
) : Request {
override val hostname: String
get() = dnsRequest.questions.single().name
Expand Down
18 changes: 2 additions & 16 deletions okhttp/api/android/okhttp.api
Original file line number Diff line number Diff line change
Expand Up @@ -1382,23 +1382,9 @@ public abstract class okhttp3/WebSocketListener {
}

public final class okhttp3/android/AndroidDns : okhttp3/Dns {
public fun <init> ()V
public fun <init> (Landroid/net/DnsResolver;Landroid/net/Network;Lokhttp3/DnsCache;ZLjava/util/concurrent/Executor;)V
public synthetic fun <init> (Landroid/net/DnsResolver;Landroid/net/Network;Lokhttp3/DnsCache;ZLjava/util/concurrent/Executor;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
public fun lookup (Ljava/lang/String;)Ljava/util/List;
public fun newCall (Lokhttp3/Dns$Request;)Lokhttp3/Dns$Call;
}

public final class okhttp3/android/EchAwareDns : okhttp3/Dns {
public static final field Companion Lokhttp3/android/EchAwareDns$Companion;
public fun <init> ()V
public fun <init> (Lokhttp3/Dns;Lokhttp3/Dns;Landroid/security/NetworkSecurityPolicy;)V
public synthetic fun <init> (Lokhttp3/Dns;Lokhttp3/Dns;Landroid/security/NetworkSecurityPolicy;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
public fun <init> (Landroid/net/DnsResolver;Landroid/net/Network;Lokhttp3/DnsCache;Z)V

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

the sea of red is nice.

public synthetic fun <init> (Landroid/net/DnsResolver;Landroid/net/Network;Lokhttp3/DnsCache;ZILkotlin/jvm/internal/DefaultConstructorMarker;)V
public fun lookup (Ljava/lang/String;)Ljava/util/List;
public fun newCall (Lokhttp3/Dns$Request;)Lokhttp3/Dns$Call;
}

public final class okhttp3/android/EchAwareDns$Companion {
public final fun forNetwork (Landroid/net/Network;)Lokhttp3/android/EchAwareDns;
}

Loading
Loading