diff --git a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsCall.kt b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsCall.kt index 76004e351f9a..c2924423001a 100644 --- a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsCall.kt +++ b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsCall.kt @@ -35,6 +35,7 @@ import okhttp3.internal.dns.DnsCallStateMachine import okhttp3.internal.dns.DnsMessage import okhttp3.internal.dns.DnsMessageReader import okhttp3.internal.dns.DnsMessageWriter +import okhttp3.internal.dns.Question import okhttp3.internal.platform.Platform import okio.Buffer import okio.BufferedSink @@ -55,8 +56,7 @@ internal class DnsOverHttpsCall( includeServiceMetadata: Boolean, canceledException: IOException?, ) : Dns.Call, - DnsCallStateMachine.Transport, - Callback { + DnsCallStateMachine.Transport { private val stateMachine = DnsCallStateMachine( transport = this, @@ -66,8 +66,8 @@ internal class DnsOverHttpsCall( includeServiceMetadata = includeServiceMetadata, ) - override fun newQuery(dnsMessage: DnsMessage): Call { - val queryParameter = dnsMessage.asQueryParameter() + override fun newQuery(question: Question): Call { + val dnsMessage = DnsMessage.query(question) return client.newCall( request = Request @@ -79,7 +79,8 @@ internal class DnsOverHttpsCall( cacheUrlOverride( dnsUrl .newBuilder() - .addQueryParameter("query", queryParameter) + .addQueryParameter("hostname", question.name) + .addQueryParameter("type", question.type.toString()) .build(), ) post(QueryRequestBody(dnsMessage)) @@ -87,7 +88,7 @@ internal class DnsOverHttpsCall( val requestUrl = dnsUrl .newBuilder() - .addQueryParameter("dns", queryParameter) + .addQueryParameter("dns", dnsMessage.asQueryParameter()) .build() url(requestUrl) } @@ -95,35 +96,40 @@ internal class DnsOverHttpsCall( ) } - override fun enqueue(query: Call) { - query.enqueue(this) + override fun enqueue( + query: Call, + callback: DnsCallStateMachine.Transport.Callback, + ) { + query.enqueue( + object : Callback { + override fun onFailure( + call: Call, + e: IOException, + ) { + callback.onFailure(e) + } + + override fun onResponse( + call: Call, + response: Response, + ) { + val dnsMessage = + try { + decodeResponse(response) + } catch (e: IOException) { + return callback.onFailure(e) + } + + callback.onResponse(dnsMessage) + } + }, + ) } override fun cancel(query: Call) { query.cancel() } - override fun onFailure( - call: Call, - e: IOException, - ) { - stateMachine.onQueryFailure(call, e) - } - - override fun onResponse( - call: Call, - response: Response, - ) { - val dnsMessage = - try { - decodeResponse(response) - } catch (e: IOException) { - return stateMachine.onQueryFailure(call, e) - } - - stateMachine.onQueryResponse(call, dnsMessage) - } - override fun enqueue(callback: Dns.Callback) { stateMachine.start(callback) } diff --git a/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/DnsRecordCodecTest.kt b/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/DnsRecordCodecTest.kt index b8279a55380e..4297f646c699 100644 --- a/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/DnsRecordCodecTest.kt +++ b/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/DnsRecordCodecTest.kt @@ -25,6 +25,7 @@ import kotlin.test.assertFailsWith import okhttp3.dnsoverhttps.internal.asQueryParameter import okhttp3.internal.dns.DnsMessage import okhttp3.internal.dns.DnsMessageReader +import okhttp3.internal.dns.Question import okhttp3.internal.dns.RESPONSE_CODE_SUCCESS import okhttp3.internal.dns.ResourceRecord import okhttp3.internal.dns.TYPE_A @@ -44,7 +45,7 @@ class DnsRecordCodecTest { private fun encodeQuery( host: String, type: Int, - ): String = DnsMessage.query(host, type).asQueryParameter() + ): String = DnsMessage.query(Question(host, type)).asQueryParameter() @Test fun testGoogleDotComEncodingWithIPv6() { diff --git a/okhttp-testing-support/src/main/kotlin/okhttp3/internal/concurrent/TaskFaker.kt b/okhttp-testing-support/src/main/kotlin/okhttp3/internal/concurrent/TaskFaker.kt index 7cad145854ee..7c018b426062 100644 --- a/okhttp-testing-support/src/main/kotlin/okhttp3/internal/concurrent/TaskFaker.kt +++ b/okhttp-testing-support/src/main/kotlin/okhttp3/internal/concurrent/TaskFaker.kt @@ -19,6 +19,7 @@ "INVISIBLE_MEMBER", "INVISIBLE_REFERENCE", ) +@file:OptIn(ExperimentalTime::class) package okhttp3.internal.concurrent @@ -30,6 +31,9 @@ import java.util.concurrent.BlockingQueue import java.util.concurrent.Executors import java.util.concurrent.TimeUnit import java.util.logging.Logger +import kotlin.time.AbstractLongTimeSource +import kotlin.time.DurationUnit +import kotlin.time.ExperimentalTime import okhttp3.TestUtil.threadFactory /** @@ -82,6 +86,12 @@ class TaskFaker : Closeable { /** Guarded by `this`. */ private var activeThreads = 0 + /** Adapt this API to Kotlin's time API. */ + val timeSource = + object : AbstractLongTimeSource(DurationUnit.NANOSECONDS) { + override fun read() = nanoTime + } + /** A task runner that posts tasks to this fake. Tasks won't be executed until requested. */ val taskRunner: TaskRunner = TaskRunner( diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt index d8858863c64a..267f73585fe1 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt @@ -34,8 +34,8 @@ import okhttp3.internal.OkHttpInternalApi * * A few things conspire to make concurrency tricky: * - * * Each DNS record type is queried in parallel; [onQueryResponse] and [onQueryFailure] may be - * called concurrently. + * * Each DNS record type is queried in parallel; [Transport.Callback.onResponse] and + * [Transport.Callback.onFailure] may be called concurrently. * * Calls to [okhttp3.Dns.Callback] must be serialized. * * We don't want to use locks to guard access to [okhttp3.Dns.Callback] functions. * @@ -66,20 +66,20 @@ class DnsCallStateMachine( get() = state.get().canceled fun start(callback: Dns.Callback) { - val queryMessages = + val questions = buildList { if (includeServiceMetadata) { - add(DnsMessage.query(call.request.hostname, TYPE_HTTPS)) + add(Question(call.request.hostname, TYPE_HTTPS)) } if (includeIPv6) { - add(DnsMessage.query(call.request.hostname, TYPE_AAAA)) + add(Question(call.request.hostname, TYPE_AAAA)) } - add(DnsMessage.query(call.request.hostname, TYPE_A)) + add(Question(call.request.hostname, TYPE_A)) } val queries = - queryMessages.map { dnsMessage -> - transport.newQuery(dnsMessage) + questions.map { question -> + transport.newQuery(question) } while (true) { @@ -100,7 +100,26 @@ class DnsCallStateMachine( if (previous.canceled || canceledException != null) { transport.cancel(query) } - transport.enqueue(query) + + transport.enqueue( + query = query, + callback = + object : Transport.Callback { + override fun onResponse(dnsResponse: DnsMessage) { + updateStateAndCallCallbacks( + completedQuery = query, + dnsResponse = dnsResponse, + ) + } + + override fun onFailure(e: IOException) { + updateStateAndCallCallbacks( + completedQuery = query, + newException = e, + ) + } + }, + ) } return @@ -122,18 +141,8 @@ class DnsCallStateMachine( } } - fun onQueryFailure( - query: Q, - e: IOException, - ) { - updateStateAndCallCallbacks( - completedQuery = query, - newException = e, - ) - } - - fun onQueryResponse( - query: Q, + private fun updateStateAndCallCallbacks( + completedQuery: Q, dnsResponse: DnsMessage, ) { val resourceRecords = @@ -145,7 +154,7 @@ class DnsCallStateMachine( } } catch (e: IOException) { return updateStateAndCallCallbacks( - completedQuery = query, + completedQuery = completedQuery, newException = e, ) } @@ -180,7 +189,7 @@ class DnsCallStateMachine( } updateStateAndCallCallbacks( - completedQuery = query, + completedQuery = completedQuery, newRecords = dnsRecords, ) } @@ -322,11 +331,20 @@ class DnsCallStateMachine( } interface Transport { - fun newQuery(dnsMessage: DnsMessage): Q + fun newQuery(question: Question): Q - fun enqueue(query: Q) + fun enqueue( + query: Q, + callback: Callback, + ) fun cancel(query: Q) + + interface Callback { + fun onFailure(e: IOException) + + fun onResponse(dnsResponse: DnsMessage) + } } } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt index df84b9fe4404..a6e16739fc11 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt @@ -34,10 +34,7 @@ data class DnsMessage( get() = (flags and 0b0000_0000_0000_1111) companion object { - fun query( - hostname: String, - type: Int, - ): DnsMessage { + fun query(question: Question): DnsMessage { // QR = 0 (Query) // RD = 1 (Recursion Desired) // OPCODE = 0 (standard query) @@ -46,24 +43,20 @@ data class DnsMessage( return DnsMessage( id = 0, flags = flags, - questions = - listOf( - Question( - name = hostname, - type = type, - ), - ), + questions = listOf(question), ) } } } +@OkHttpInternalApi data class Question( val name: String, val type: Int, val `class`: Int = CLASS_IN, ) +@OkHttpInternalApi sealed interface ResourceRecord { val name: String val timeToLive: Int diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt new file mode 100644 index 000000000000..1ad0c04f3d5a --- /dev/null +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt @@ -0,0 +1,285 @@ +/* + * Copyright (c) 2026 OkHttp Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package okhttp3.internal.dns + +import java.io.IOException +import java.util.concurrent.ConcurrentHashMap +import java.util.concurrent.atomic.AtomicReference +import kotlin.time.ComparableTimeMark as Time +import kotlin.time.Duration +import kotlin.time.Duration.Companion.seconds +import kotlin.time.ExperimentalTime +import kotlin.time.TimeSource +import okhttp3.internal.OkHttpInternalApi +import okhttp3.internal.concurrent.TaskRunner +import okhttp3.internal.dns.DnsCallStateMachine.Transport + +// TODO: evict old entries from cache using State.lastRequestedAt + +/** + * A DNS transport that caches responses according to their [ResourceRecord.timeToLive], bounded by + * a user-supplied minimum and maximum cache duration. + * + * The age of the result impacts how queries are satisfied: + * + * * After [Result.expireAt], the cached result is not used and a call to the underlying transport + * is made. + * + * * After [Result.revalidateAt], the cached result is returned immediately. A call to the + * underlying transport is also made, in order to freshen the cache for a possible future call. + * + * * Otherwise, the cached data is returned immediately. + * + * Failures are cached to prevent error cases from using more resources than success cases. There's + * no server-provided defaults for these so the configuration parameter [failureTimeToLive] must be + * used. + * + * If this receives multiple equivalent queries, it combines them into a single query on the + * underlying transport. + */ +@OkHttpInternalApi +@OptIn(ExperimentalTime::class) // We know Clock and Instant will be stable in Kotlin 2.3. +class CachingTransport( + private val taskRunner: TaskRunner, + private val delegate: Transport, + private val timeSource: TimeSource.WithComparableMarks, + private val minimumTimeToLive: Duration = 10.seconds, + private val maximumTimeToLive: Duration = 300.seconds, + private val failureTimeToLive: Duration = 10.seconds, + private val revalidateBeforeExpire: Duration = 5.seconds, +) : Transport> { + private val entries = ConcurrentHashMap() + + init { + require(failureTimeToLive >= 0.seconds) + require(minimumTimeToLive >= 0.seconds) + require(maximumTimeToLive >= minimumTimeToLive) + require(revalidateBeforeExpire >= 0.seconds) + } + + override fun newQuery(question: Question): Query { + val inserted = Entry(question) + val entry = entries.putIfAbsent(question, inserted) ?: inserted + return Query(entry) + } + + override fun enqueue( + query: Query, + callback: Transport.Callback>, + ) { + check(query.callback == null) { "already enqueued" } + query.callback = callback + + val entry = query.entry + val now = timeSource.markNow() + while (true) { + val previous = entry.state.get() + val result = previous.result + val inFlightCall = previous.inFlightCall + + // We use a cached value unless it's expired. + val useCached = result != null && now < result.expireAt + + // Revalidate the cache if necessary. Note that we might revalidate the cache without any + // particular callback waiting for that response. + val next = + previous.copy( + lastRequestedAt = now, + inFlightCall = + when { + inFlightCall != null && useCached -> { + inFlightCall + } + + inFlightCall != null -> { + inFlightCall.copy(queries = inFlightCall.queries + query) + } + + result == null || now >= result.revalidateAt -> { + InFlightCall( + query = delegate.newQuery(entry.question), + sentAt = now, + queries = if (useCached) listOf() else listOf(query), + ) + } + + else -> { + null + } + }, + ) + + if (!entry.state.compareAndSet(previous, next)) continue // Lost a race, retry. + + if (inFlightCall == null && next.inFlightCall != null) { + delegate.enqueue(next.inFlightCall.query, entry) + } + + if (useCached) { + taskRunner.newQueue().execute("${query.entry.question.name} dns") { + when (result) { + is Result.Success -> callback.onResponse(result.message) + is Result.Failure -> callback.onFailure(result.exception) + } + } + } + + return + } + } + + /** + * Note that we don't cancel the query even if nothing is waiting on it. We assume there's still + * value in updating the cache! + */ + override fun cancel(query: Query) { + while (true) { + val entry = query.entry + val previous = entry.state.get() + val inFlightCall = previous.inFlightCall ?: return + + // If we've already called the callback, there's nothing to do. + val newQueries = inFlightCall.queries - query + if (newQueries.size == inFlightCall.queries.size) return + + val next = + previous.copy( + inFlightCall = + inFlightCall.copy( + queries = newQueries, + ), + ) + + if (!entry.state.compareAndSet(previous, next)) continue // Lost a race, retry. + + taskRunner.newQueue().execute("${query.entry.question.name} dns") { + query.callback!!.onFailure(IOException("canceled")) + } + + return + } + } + + /** A query on this transport. */ + class Query( + val entry: CachingTransport.Entry, + ) { + var callback: Transport.Callback>? = null + } + + /** + * Transforms a series of queries on this transport to a smaller (or at least not larger) series + * of queries on the underlying transport. + */ + inner class Entry( + val question: Question, + ) : Transport.Callback { + val state = AtomicReference(State()) + + override fun onFailure(e: IOException) { + while (true) { + val previous = state.get() + val sentAt = previous.inFlightCall!!.sentAt + val revalidateDelay = (failureTimeToLive - revalidateBeforeExpire).coerceAtLeast(0.seconds) + + val next = + previous.copy( + inFlightCall = null, + result = + Result.Failure( + exception = e, + revalidateAt = sentAt + revalidateDelay, + expireAt = sentAt + failureTimeToLive, + ), + ) + + if (!state.compareAndSet(previous, next)) continue // Lost a race, retry. + + val queries = previous.inFlightCall.queries + for (query in queries) { + query.callback!!.onFailure(e) + } + + return + } + } + + override fun onResponse(dnsResponse: DnsMessage) { + while (true) { + val previous = state.get() + val sentAt = previous.inFlightCall!!.sentAt + val timeToLive = + (dnsResponse.answers.minOfOrNull { it.timeToLive } ?: 0) + .seconds + .coerceIn(minimumTimeToLive, maximumTimeToLive) + val revalidateDelay = (timeToLive - revalidateBeforeExpire).coerceAtLeast(0.seconds) + + val next = + previous.copy( + inFlightCall = null, + result = + Result.Success( + message = dnsResponse, + revalidateAt = sentAt + revalidateDelay, + expireAt = sentAt + timeToLive, + ), + ) + + if (!state.compareAndSet(previous, next)) continue // Lost a race, retry. + + val queries = previous.inFlightCall.queries + for (query in queries) { + query.callback!!.onResponse(dnsResponse) + } + + return + } + } + } + + /** A snapshot of the state of a single entry. */ + data class State( + val lastRequestedAt: Time? = null, + val inFlightCall: InFlightCall? = null, + val result: Result? = null, + ) + + /** A call to the underlying transport. */ + data class InFlightCall( + val query: Q, + val sentAt: Time, + /** The possibly-empty set of queries to notify when this call is complete. */ + val queries: List>, + ) + + /** A cached result. */ + sealed interface Result { + val revalidateAt: Time + val expireAt: Time + + class Failure( + override val revalidateAt: Time, + override val expireAt: Time, + val exception: IOException, + ) : Result + + class Success( + override val revalidateAt: Time, + override val expireAt: Time, + val message: DnsMessage, + ) : Result + } +} diff --git a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt index f123821707fe..6a5b162afd81 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt @@ -17,84 +17,218 @@ package okhttp3.internal.dns -import java.io.IOException +import app.cash.burst.Burst import java.net.InetAddress import kotlin.test.Test import okhttp3.Dns import okhttp3.Protocol import okhttp3.internal.OkHttpInternalApi +@Burst class DnsCallStateMachineTest { + /** Arbitrary sample values. */ + private val blueIpv6s = listOf(InetAddress.getByName("1:2::3:4")) + private val blueIpv4s = listOf(InetAddress.getByName("10.20.30.40")) + @Test - fun `happy path`() = - testDnsCallStateMachine( - request = Dns.Request(hostname = "lysine.dev"), - ) { - enqueue() + fun `happy path`(caching: Boolean = true) { + testDnsCallStateMachine { + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = caching, + ) + call.enqueue() - val query0 = takeQuery("lysine.dev", TYPE_HTTPS) - val query1 = takeQuery("lysine.dev", TYPE_AAAA) - val query2 = takeQuery("lysine.dev", TYPE_A) + val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = transport.takeQuery("lysine.dev", TYPE_A) - respondIpAddresses( - query = query1.query, - addresses = listOf(InetAddress.getByName("1:2::3:4")), + query1.respondIpAddresses( + addresses = blueIpv6s, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("1:2::3:4")), + call.takeOnRecordsIpAddresses( + addresses = blueIpv6s, ) - respondIpAddresses( - query = query2.query, - addresses = listOf(InetAddress.getByName("10.20.30.40")), + query2.respondIpAddresses( + addresses = blueIpv4s, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + call.takeOnRecordsIpAddresses( + addresses = blueIpv4s, ) - respondServiceMetadata( - query = query0.query, + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) - takeOnRecordsServiceMetadata( + call.takeOnRecordsServiceMetadata( last = true, alpnIds = listOf(Protocol.HTTP_2), ) } + } + + @Test + fun `cache already completed values`() = + testDnsCallStateMachine { + val call0 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + includeServiceMetadata = false, + caching = true, + ) + call0.enqueue() + + val call0QueryIpv6 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val call0QueryIpv4 = transport.takeQuery("lysine.dev", TYPE_A) + + call0QueryIpv6.respondIpAddresses( + addresses = blueIpv6s, + ) + call0.takeOnRecordsIpAddresses( + addresses = blueIpv6s, + ) + + call0QueryIpv4.respondIpAddresses( + addresses = blueIpv4s, + ) + call0.takeOnRecordsIpAddresses( + last = true, + addresses = blueIpv4s, + ) + + val call1 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + includeServiceMetadata = false, + caching = true, + ) + call1.enqueue() + + call1.takeOnRecordsIpAddresses( + addresses = blueIpv6s, + ) + call1.takeOnRecordsIpAddresses( + last = true, + addresses = blueIpv4s, + ) + } + /** Confirm that two queries to the cache yield a single query to the underlying transport. */ @Test - fun `failure returned last`() = - testDnsCallStateMachine( - request = Dns.Request(hostname = "lysine.dev"), - ) { - enqueue() + fun `cache in flight calls`() = + testDnsCallStateMachine { + val call0 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + includeServiceMetadata = false, + caching = true, + ) + call0.enqueue() - val query0 = takeQuery("lysine.dev", TYPE_HTTPS) - val query1 = takeQuery("lysine.dev", TYPE_AAAA) - val query2 = takeQuery("lysine.dev", TYPE_A) + val call0QueryIpv6 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val call0QueryIpv4 = transport.takeQuery("lysine.dev", TYPE_A) - respondFailure( - query = query1.query, - e = IOException("boom!"), + val call1 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + includeServiceMetadata = false, + caching = true, + ) + call1.enqueue() + + call0QueryIpv6.respondIpAddresses( + addresses = blueIpv6s, + ) + call0.takeOnRecordsIpAddresses( + addresses = blueIpv6s, + ) + call1.takeOnRecordsIpAddresses( + addresses = blueIpv6s, ) - respondIpAddresses( - query = query2.query, - addresses = listOf(InetAddress.getByName("10.20.30.40")), + call0QueryIpv4.respondIpAddresses( + addresses = blueIpv4s, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + call0.takeOnRecordsIpAddresses( + last = true, + addresses = blueIpv4s, ) + call1.takeOnRecordsIpAddresses( + last = true, + addresses = blueIpv4s, + ) + } - respondServiceMetadata( - query = query0.query, + @Test + fun `failure returned last`(caching: Boolean = true) = + testDnsCallStateMachine { + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = caching, + ) + call.enqueue() + + val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = transport.takeQuery("lysine.dev", TYPE_A) + + query1.respondFailure("boom!") + + query2.respondIpAddresses( + addresses = blueIpv4s, + ) + call.takeOnRecordsIpAddresses( + addresses = blueIpv4s, + ) + + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) - takeOnRecordsServiceMetadata( + call.takeOnRecordsServiceMetadata( alpnIds = listOf(Protocol.HTTP_2), ) - takeOnFailure("boom!") + call.takeOnFailure("boom!") + } + + @Test + fun `failure is cached`() = + testDnsCallStateMachine { + val call0 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + includeServiceMetadata = false, + ) + call0.enqueue() + + val queryIpv6 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val queryIpv4 = transport.takeQuery("lysine.dev", TYPE_A) + + queryIpv6.respondFailure("boom!") + queryIpv4.respondIpAddresses( + addresses = blueIpv4s, + ) + + call0.takeOnRecordsIpAddresses( + addresses = blueIpv4s, + ) + call0.takeOnFailure("boom!") + + val call1 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + includeServiceMetadata = false, + ) + call1.enqueue() + + call1.takeOnRecordsIpAddresses( + addresses = blueIpv4s, + ) + call1.takeOnFailure("boom!") } /** @@ -105,32 +239,34 @@ class DnsCallStateMachineTest { * re-entrant call on a single thread. */ @Test - fun `calls to onRecords are serialized`() = - testDnsCallStateMachine(request = Dns.Request(hostname = "lysine.dev")) { - enqueue() + fun `calls to onRecords are serialized`(caching: Boolean = true) = + testDnsCallStateMachine { + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = caching, + ) + call.enqueue() - val query0 = takeQuery("lysine.dev", TYPE_HTTPS) - val query1 = takeQuery("lysine.dev", TYPE_AAAA) - val query2 = takeQuery("lysine.dev", TYPE_A) + val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = transport.takeQuery("lysine.dev", TYPE_A) onNextEvent = { - respondIpAddresses( - query = query2.query, - addresses = listOf(InetAddress.getByName("10.20.30.40")), + query2.respondIpAddresses( + addresses = blueIpv4s, ) - respondIpAddresses( - query = query1.query, - addresses = listOf(InetAddress.getByName("1:2::3:4")), + query1.respondIpAddresses( + addresses = blueIpv6s, ) } - respondServiceMetadata( - query = query0.query, + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) - takeOnRecordsServiceMetadata( + call.takeOnRecordsServiceMetadata( alpnIds = listOf(Protocol.HTTP_2), ) - takeOnRecordsIpAddresses( + call.takeOnRecordsIpAddresses( last = true, addresses = listOf( @@ -146,66 +282,125 @@ class DnsCallStateMachineTest { */ @Test fun `cancel before enqueue`() = - testDnsCallStateMachine( - request = Dns.Request(hostname = "lysine.dev"), - ) { + testDnsCallStateMachine { + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = false, + ) + call.cancel() + call.enqueue() + + transport.takeCancel("lysine.dev", TYPE_HTTPS) + val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) + transport.takeCancel("lysine.dev", TYPE_AAAA) + val query1 = transport.takeQuery("lysine.dev", TYPE_AAAA) + transport.takeCancel("lysine.dev", TYPE_A) + val query2 = transport.takeQuery("lysine.dev", TYPE_A) + + query0.respondFailure("canceled") + query1.respondFailure("canceled") + query2.respondFailure("canceled") + + call.takeOnFailure("canceled") + } + + @Test + fun `cancel before enqueue with caching`() = + testDnsCallStateMachine { + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + ) call.cancel() - enqueue() + call.enqueue() - val query0 = takeCancel("lysine.dev", TYPE_HTTPS) - takeQuery("lysine.dev", TYPE_HTTPS) - val query1 = takeCancel("lysine.dev", TYPE_AAAA) - takeQuery("lysine.dev", TYPE_AAAA) - val query2 = takeCancel("lysine.dev", TYPE_A) - takeQuery("lysine.dev", TYPE_A) + val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = transport.takeQuery("lysine.dev", TYPE_A) - respondFailure(query0.query, IOException("canceled")) - respondFailure(query1.query, IOException("canceled")) - respondFailure(query2.query, IOException("canceled")) + query0.respondFailure("canceled") + query1.respondFailure("canceled") + query2.respondFailure("canceled") - takeOnFailure("canceled") + call.takeOnFailure("canceled") } /** Cancels are asynchronous and if the canceled query completes anyway, that's fine. */ @Test fun `cancel ignored if canceled query completes`() = - testDnsCallStateMachine( - request = Dns.Request(hostname = "lysine.dev"), - ) { - enqueue() + testDnsCallStateMachine { + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = false, + ) + call.enqueue() - val query0 = takeQuery("lysine.dev", TYPE_HTTPS) - val query1 = takeQuery("lysine.dev", TYPE_AAAA) - val query2 = takeQuery("lysine.dev", TYPE_A) + val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = transport.takeQuery("lysine.dev", TYPE_A) - respondIpAddresses( - query = query1.query, - addresses = listOf(InetAddress.getByName("1:2::3:4")), + query1.respondIpAddresses( + addresses = blueIpv6s, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("1:2::3:4")), + call.takeOnRecordsIpAddresses( + addresses = blueIpv6s, ) call.cancel() - takeCancel("lysine.dev", TYPE_HTTPS) - takeCancel("lysine.dev", TYPE_A) + transport.takeCancel("lysine.dev", TYPE_HTTPS) + transport.takeCancel("lysine.dev", TYPE_A) - respondIpAddresses( - query = query2.query, - addresses = listOf(InetAddress.getByName("10.20.30.40")), + query2.respondIpAddresses( + addresses = blueIpv4s, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + call.takeOnRecordsIpAddresses( + addresses = blueIpv4s, ) - respondServiceMetadata( - query = query0.query, + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) - takeOnRecordsServiceMetadata( + call.takeOnRecordsServiceMetadata( last = true, alpnIds = listOf(Protocol.HTTP_2), ) } + + /** When caching, cancels aren't applied to the transport. */ + @Test + fun `cancel ignored if canceled query completes with caching`() = + testDnsCallStateMachine { + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + ) + call.enqueue() + + val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = transport.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = transport.takeQuery("lysine.dev", TYPE_A) + + query1.respondIpAddresses( + addresses = blueIpv6s, + ) + call.takeOnRecordsIpAddresses( + addresses = blueIpv6s, + ) + + call.cancel() + + query2.respondIpAddresses( + addresses = blueIpv4s, + ) + + query0.respondServiceMetadata( + alpnIds = listOf("h2"), + ) + call.takeOnFailure("canceled") + } } diff --git a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt index 0e088aae3423..0d122444c4a7 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt @@ -1,19 +1,26 @@ -@file:OptIn(OkHttpInternalApi::class) +@file:OptIn(OkHttpInternalApi::class, ExperimentalTime::class) package okhttp3.internal.dns import assertk.assertThat import assertk.assertions.hasMessage import assertk.assertions.isEqualTo +import assertk.assertions.isNull import java.io.IOException import java.net.InetAddress import java.util.concurrent.LinkedBlockingDeque +import kotlin.time.ExperimentalTime import okhttp3.Dns import okhttp3.Protocol import okhttp3.dnsResponse import okhttp3.internal.OkHttpInternalApi -import okhttp3.internal.dns.DnsCallStateMachineTester.Event.OnRecords -import okhttp3.internal.dns.DnsCallStateMachineTester.Event.QueryEnqueued +import okhttp3.internal.concurrent.TaskFaker +import okhttp3.internal.dns.DnsCallStateMachine.Transport +import okhttp3.internal.dns.DnsCallStateMachineTester.CallEvent.OnFailure +import okhttp3.internal.dns.DnsCallStateMachineTester.CallEvent.OnRecords +import okhttp3.internal.dns.DnsCallStateMachineTester.TransportEvent.QueryCanceled +import okhttp3.internal.dns.DnsCallStateMachineTester.TransportEvent.QueryEnqueued +import okhttp3.internal.dns.DnsMessage.Companion.query import okio.ByteString /** @@ -24,278 +31,297 @@ import okio.ByteString * This tracks all effects from the state machine as events: creating queries, canceling queries, * calling callbacks. */ -fun testDnsCallStateMachine( - request: Dns.Request, - includeIPv6: Boolean = true, - includeServiceMetadata: Boolean = true, - block: DnsCallStateMachineTester.() -> Unit, -) { - val tester = DnsCallStateMachineTester(request, includeIPv6, includeServiceMetadata) +fun testDnsCallStateMachine(block: DnsCallStateMachineTester.() -> Unit) { + val tester = DnsCallStateMachineTester() tester.block() + assertThat(tester.transport.events.poll(), "unexpected transport event").isNull() } -class DnsCallStateMachineTester internal constructor( - request: Dns.Request, - includeIPv6: Boolean = true, - includeServiceMetadata: Boolean = true, -) { - private val events = LinkedBlockingDeque() +class DnsCallStateMachineTester internal constructor() { var onNextEvent: (() -> Unit)? = null /** Defend against re-entrant calls. */ private var acceptCallbacks: Boolean = true - private val transport = - object : DnsCallStateMachine.Transport { - override fun newQuery(dnsMessage: DnsMessage) = Query(dnsMessage) + val transport = Transport() - override fun enqueue(query: Query) { - postEvent(QueryEnqueued(query)) - } + private val taskFaker = TaskFaker() - override fun cancel(query: Query) { - postEvent(Event.QueryCanceled(query)) - } - } + private val cachingTransport = + CachingTransport( + taskRunner = taskFaker.taskRunner, + delegate = transport, + timeSource = taskFaker.timeSource, + ) - val call: Dns.Call = - object : Dns.Call { - override val request: Dns.Request = request + fun newCall( + request: Dns.Request, + includeIPv6: Boolean = true, + includeServiceMetadata: Boolean = true, + caching: Boolean = false, + ): Call = + Call( + request = request, + includeIPv6 = includeIPv6, + includeServiceMetadata = includeServiceMetadata, + caching = caching, + ) - override fun enqueue(callback: Dns.Callback) { - stateMachine.start(callback) - } + inner class Transport : DnsCallStateMachine.Transport { + val events = LinkedBlockingDeque() - override fun cancel() { - stateMachine.cancel() - } + override fun newQuery(question: Question) = Query(question) + + private fun postEvent(e: TransportEvent) { + events.put(e) - override fun isCanceled() = stateMachine.canceled + onNextEvent?.invoke() + onNextEvent = null } - private val callback = - object : Dns.Callback { - override fun onRecords( - call: Dns.Call, - last: Boolean, - records: List, - ) { - check(call == this@DnsCallStateMachineTester.call) - check(acceptCallbacks) { "unexpected callback" } - - acceptCallbacks = false - try { - postEvent(OnRecords(last, records)) - } finally { - acceptCallbacks = true - } - } + private fun takeEvent(): TransportEvent { + taskFaker.runTasks() // Run any queued async work first. + return events.take() + } - override fun onFailure( - call: Dns.Call, - e: IOException, - ) { - check(call == this@DnsCallStateMachineTester.call) - check(acceptCallbacks) { "unexpected callback" } - - acceptCallbacks = false - try { - postEvent(Event.OnFailure(e)) - } finally { - acceptCallbacks = true - } - } + /** Asserts that the next-posted event is a query enqueue. */ + fun takeQuery( + hostname: String, + type: Int, + ): QueryEnqueued { + val event = transport.takeEvent() as QueryEnqueued + assertThat(event.hostname).isEqualTo(hostname) + assertThat(event.type).isEqualTo(type) + return event } - val stateMachine = - DnsCallStateMachine( - transport = transport, - call = call, - canceledException = null, - includeIPv6 = includeIPv6, - includeServiceMetadata = includeServiceMetadata, - ) + /** Asserts that the next-posted event is a query cancel. */ + fun takeCancel( + hostname: String, + type: Int, + ): QueryCanceled { + val event = transport.takeEvent() as QueryCanceled + assertThat(event.hostname).isEqualTo(hostname) + assertThat(event.type).isEqualTo(type) + return event + } - /** Start the DNS call. */ - fun enqueue() { - check(acceptCallbacks) { "unexpected enqueue" } + override fun enqueue( + query: Query, + callback: Transport.Callback, + ) { + postEvent(QueryEnqueued(query, callback)) + } - acceptCallbacks = false - try { - call.enqueue(callback) - } finally { - acceptCallbacks = true + override fun cancel(query: Query) { + postEvent(QueryCanceled(query)) } } - private fun postEvent(e: Event) { - events.put(e) + /** A DNS call for the fake state machine. */ + inner class Call( + override val request: Dns.Request, + includeIPv6: Boolean = true, + includeServiceMetadata: Boolean = true, + caching: Boolean = false, + ) : Dns.Call, + Dns.Callback { + private val events = LinkedBlockingDeque() + + val stateMachine = + DnsCallStateMachine( + transport = + when { + caching -> cachingTransport + else -> transport + }, + call = this, + canceledException = null, + includeIPv6 = includeIPv6, + includeServiceMetadata = includeServiceMetadata, + ) + + fun enqueue() { + check(acceptCallbacks) { "unexpected enqueue" } + acceptCallbacks = false + try { + enqueue(this) + } finally { + acceptCallbacks = true + } + } - onNextEvent?.invoke() - onNextEvent = null - } + override fun enqueue(callback: Dns.Callback) { + stateMachine.start(callback) + } - /** Asserts that the next-posted event is a query enqueue. */ - fun takeQuery( - hostname: String, - type: Int, - ): QueryEnqueued { - val event = events.take() as QueryEnqueued - assertThat(event.hostname).isEqualTo(hostname) - assertThat(event.type).isEqualTo(type) - return event - } + override fun cancel() { + stateMachine.cancel() + } - /** Asserts that the next-posted event is a query cancel. */ - fun takeCancel( - hostname: String, - type: Int, - ): Event.QueryCanceled { - val event = events.take() as Event.QueryCanceled - assertThat(event.hostname).isEqualTo(hostname) - assertThat(event.type).isEqualTo(type) - return event - } + override fun isCanceled() = stateMachine.canceled - /** Respond to a [TYPE_A] or [TYPE_AAAA] query with a (possibly-empty) list of IP addresses. */ - fun respondIpAddresses( - query: Query, - timeToLive: Int = 300, - addresses: List = listOf(), - ) { - stateMachine.onQueryResponse( - query, - dnsResponse( - query.dnsMessage, - addresses.map { address -> - ResourceRecord.IpAddress( - name = - query.dnsMessage.questions - .single() - .name, - timeToLive = timeToLive, - address = address, - ) - }, - ), - ) - } + private fun postEvent(e: CallEvent) { + events.put(e) - /** Respond to a [TYPE_HTTPS] query with service metadata. */ - fun respondServiceMetadata( - query: Query, - timeToLive: Int = 300, - alpnIds: List? = null, - echConfigList: ByteString? = null, - ) { - stateMachine.onQueryResponse( - query, - dnsResponse( - query.dnsMessage, - listOf( - ResourceRecord.Https( - name = - query.dnsMessage.questions - .single() - .name, - timeToLive = timeToLive, - alpnIds = alpnIds, - echConfigList = echConfigList, - ), - ), - ), - ) - } + onNextEvent?.invoke() + onNextEvent = null + } - /** Respond to any query with a failure. */ - fun respondFailure( - query: Query, - e: IOException, - ) { - stateMachine.onQueryFailure(query, e) - } + private fun takeEvent(): CallEvent { + taskFaker.runTasks() // Run any queued async work first. + return events.take() + } - /** - * Asserts that the next-posted event is a call to [Dns.Callback.onRecords] with a list of IP - * addresses. - */ - fun takeOnRecordsIpAddresses( - last: Boolean = false, - addresses: List, - ): OnRecords { - val event = events.take() as OnRecords - assertThat(event.last).isEqualTo(last) - assertThat(event.records.map { (it as Dns.Record.IpAddress).address }) - .isEqualTo(addresses) - return event - } + override fun onRecords( + call: Dns.Call, + last: Boolean, + records: List, + ) { + check(call == this) + check(acceptCallbacks) { "unexpected callback" } + + acceptCallbacks = false + try { + postEvent(OnRecords(last, records)) + } finally { + acceptCallbacks = true + } + } - /** - * Asserts that the next-posted event is a call to [Dns.Callback.onRecords] with service metadata. - */ - fun takeOnRecordsServiceMetadata( - last: Boolean = false, - alpnIds: List? = null, - echConfigList: ByteString? = null, - ): OnRecords { - val event = events.take() as OnRecords - assertThat(event.last).isEqualTo(last) - - val serviceMetadata = event.records.single() as Dns.Record.ServiceMetadata - assertThat(serviceMetadata.alpnIds).isEqualTo(alpnIds) - assertThat(serviceMetadata.echConfigList).isEqualTo(echConfigList) - return event - } + override fun onFailure( + call: Dns.Call, + e: IOException, + ) { + check(call == this) + check(acceptCallbacks) { "unexpected callback" } + + acceptCallbacks = false + try { + postEvent(OnFailure(e)) + } finally { + acceptCallbacks = true + } + } + + /** + * Asserts that the next-posted event is a call to [Dns.Callback.onRecords] with a list of IP + * addresses. + */ + fun takeOnRecordsIpAddresses( + last: Boolean = false, + addresses: List, + ): OnRecords { + val event = takeEvent() as OnRecords + assertThat(event.last).isEqualTo(last) + assertThat(event.records.map { (it as Dns.Record.IpAddress).address }) + .isEqualTo(addresses) + return event + } - /** Asserts that the next-posted event is a call to [Dns.Callback.onFailure]. */ - fun takeOnFailure(message: String): Event.OnFailure { - val event = events.take() as Event.OnFailure - assertThat(event.e).hasMessage(message) - return event + /** + * Asserts that the next-posted event is a call to [Dns.Callback.onRecords] with service metadata. + */ + fun takeOnRecordsServiceMetadata( + last: Boolean = false, + alpnIds: List? = null, + echConfigList: ByteString? = null, + ): OnRecords { + val event = takeEvent() as OnRecords + assertThat(event.last).isEqualTo(last) + + val serviceMetadata = event.records.single() as Dns.Record.ServiceMetadata + assertThat(serviceMetadata.alpnIds).isEqualTo(alpnIds) + assertThat(serviceMetadata.echConfigList).isEqualTo(echConfigList) + return event + } + + /** Asserts that the next-posted event is a call to [Dns.Callback.onFailure]. */ + fun takeOnFailure(message: String): OnFailure { + val event = takeEvent() as OnFailure + assertThat(event.e).hasMessage(message) + return event + } } class Query( - val dnsMessage: DnsMessage, + val question: Question, ) - sealed interface Event { - data class QueryEnqueued( + sealed interface TransportEvent { + class QueryEnqueued( val query: Query, - ) : Event { + val callback: Transport.Callback, + ) : TransportEvent { val hostname: String - get() = - query.dnsMessage.questions - .single() - .name + get() = query.question.name val type: Int - get() = - query.dnsMessage.questions - .single() - .type + get() = query.question.type + + /** Respond to a [TYPE_HTTPS] query with service metadata. */ + fun respondServiceMetadata( + timeToLive: Int = 300, + alpnIds: List? = null, + echConfigList: ByteString? = null, + ) { + callback.onResponse( + dnsResponse( + query(query.question), + listOf( + ResourceRecord.Https( + name = query.question.name, + timeToLive = timeToLive, + alpnIds = alpnIds, + echConfigList = echConfigList, + ), + ), + ), + ) + } + + /** Respond to any query with a failure. */ + fun respondFailure(message: String) { + callback.onFailure(IOException(message)) + } + + /** Respond to a [TYPE_A] or [TYPE_AAAA] query with a (possibly-empty) list of IP addresses. */ + fun respondIpAddresses( + timeToLive: Int = 300, + addresses: List = listOf(), + ) { + callback.onResponse( + dnsResponse( + query(query.question), + addresses.map { address -> + ResourceRecord.IpAddress( + name = query.question.name, + timeToLive = timeToLive, + address = address, + ) + }, + ), + ) + } } - data class QueryCanceled( + class QueryCanceled( val query: Query, - ) : Event { + ) : TransportEvent { val hostname: String - get() = - query.dnsMessage.questions - .single() - .name + get() = query.question.name val type: Int - get() = - query.dnsMessage.questions - .single() - .type + get() = query.question.type } + } - data class OnRecords( + sealed interface CallEvent { + class OnRecords( val last: Boolean, val records: List, - ) : Event + ) : CallEvent - data class OnFailure( + class OnFailure( val e: IOException, - ) : Event + ) : CallEvent } }