From 09a06cab0e1ebe3576e9216bcd5ec3c3585271db Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Wed, 22 Jul 2026 14:44:28 -0400 Subject: [PATCH 1/5] Switch to Question as the Transport argument It makes caching easier. --- .../internal/-DnsOverHttpsCall.kt | 58 +++---- .../dnsoverhttps/DnsRecordCodecTest.kt | 3 +- .../internal/dns/-DnsCallStateMachine.kt | 63 +++++--- .../okhttp3/internal/dns/-DnsMessage.kt | 15 +- .../internal/dns/DnsCallStateMachineTest.kt | 57 +++---- .../internal/dns/DnsCallStateMachineTester.kt | 147 ++++++++---------- 6 files changed, 158 insertions(+), 185 deletions(-) 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..a293676ae54c 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,33 +96,32 @@ internal class DnsOverHttpsCall( ) } - override fun enqueue(query: Call) { - query.enqueue(this) - } - - override fun cancel(query: Call) { - query.cancel() - } - - override fun onFailure( - call: Call, - e: IOException, + override fun enqueue( + query: Call, + callback: DnsCallStateMachine.Transport.Callback ) { - stateMachine.onQueryFailure(call, e) - } + 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) + } - override fun onResponse( - call: Call, - response: Response, - ) { - val dnsMessage = - try { - decodeResponse(response) - } catch (e: IOException) { - return stateMachine.onQueryFailure(call, e) + callback.onResponse(dnsMessage) + } } + ) + } - stateMachine.onQueryResponse(call, dnsMessage) + override fun cancel(query: Call) { + query.cancel() } override fun enqueue(callback: Dns.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/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt index d8858863c64a..e6ee3fe30148 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,25 @@ 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 +140,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 +153,7 @@ class DnsCallStateMachine( } } catch (e: IOException) { return updateStateAndCallCallbacks( - completedQuery = query, + completedQuery = completedQuery, newException = e, ) } @@ -180,7 +188,7 @@ class DnsCallStateMachine( } updateStateAndCallCallbacks( - completedQuery = query, + completedQuery = completedQuery, newRecords = dnsRecords, ) } @@ -322,11 +330,16 @@ 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/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt index f123821707fe..bf3ce4371f66 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt @@ -17,7 +17,6 @@ package okhttp3.internal.dns -import java.io.IOException import java.net.InetAddress import kotlin.test.Test import okhttp3.Dns @@ -36,24 +35,21 @@ class DnsCallStateMachineTest { val query1 = takeQuery("lysine.dev", TYPE_AAAA) val query2 = takeQuery("lysine.dev", TYPE_A) - respondIpAddresses( - query = query1.query, + query1.respondIpAddresses( addresses = listOf(InetAddress.getByName("1:2::3:4")), ) takeOnRecordsIpAddresses( addresses = listOf(InetAddress.getByName("1:2::3:4")), ) - respondIpAddresses( - query = query2.query, + query2.respondIpAddresses( addresses = listOf(InetAddress.getByName("10.20.30.40")), ) takeOnRecordsIpAddresses( addresses = listOf(InetAddress.getByName("10.20.30.40")), ) - respondServiceMetadata( - query = query0.query, + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) takeOnRecordsServiceMetadata( @@ -73,21 +69,16 @@ class DnsCallStateMachineTest { val query1 = takeQuery("lysine.dev", TYPE_AAAA) val query2 = takeQuery("lysine.dev", TYPE_A) - respondFailure( - query = query1.query, - e = IOException("boom!"), - ) + query1.respondFailure("boom!") - respondIpAddresses( - query = query2.query, + query2.respondIpAddresses( addresses = listOf(InetAddress.getByName("10.20.30.40")), ) takeOnRecordsIpAddresses( addresses = listOf(InetAddress.getByName("10.20.30.40")), ) - respondServiceMetadata( - query = query0.query, + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) takeOnRecordsServiceMetadata( @@ -114,17 +105,14 @@ class DnsCallStateMachineTest { val query2 = takeQuery("lysine.dev", TYPE_A) onNextEvent = { - respondIpAddresses( - query = query2.query, + query2.respondIpAddresses( addresses = listOf(InetAddress.getByName("10.20.30.40")), ) - respondIpAddresses( - query = query1.query, + query1.respondIpAddresses( addresses = listOf(InetAddress.getByName("1:2::3:4")), ) } - respondServiceMetadata( - query = query0.query, + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) takeOnRecordsServiceMetadata( @@ -152,16 +140,16 @@ class DnsCallStateMachineTest { call.cancel() 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) + takeCancel("lysine.dev", TYPE_HTTPS) + val query0 = takeQuery("lysine.dev", TYPE_HTTPS) + takeCancel("lysine.dev", TYPE_AAAA) + val query1 = takeQuery("lysine.dev", TYPE_AAAA) + takeCancel("lysine.dev", TYPE_A) + val query2 = 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") } @@ -178,8 +166,7 @@ class DnsCallStateMachineTest { val query1 = takeQuery("lysine.dev", TYPE_AAAA) val query2 = takeQuery("lysine.dev", TYPE_A) - respondIpAddresses( - query = query1.query, + query1.respondIpAddresses( addresses = listOf(InetAddress.getByName("1:2::3:4")), ) takeOnRecordsIpAddresses( @@ -191,16 +178,14 @@ class DnsCallStateMachineTest { takeCancel("lysine.dev", TYPE_HTTPS) takeCancel("lysine.dev", TYPE_A) - respondIpAddresses( - query = query2.query, + query2.respondIpAddresses( addresses = listOf(InetAddress.getByName("10.20.30.40")), ) takeOnRecordsIpAddresses( addresses = listOf(InetAddress.getByName("10.20.30.40")), ) - respondServiceMetadata( - query = query0.query, + query0.respondServiceMetadata( alpnIds = listOf("h2"), ) takeOnRecordsServiceMetadata( diff --git a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt index 0e088aae3423..176c45ed32db 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt @@ -12,8 +12,10 @@ import okhttp3.Dns import okhttp3.Protocol import okhttp3.dnsResponse import okhttp3.internal.OkHttpInternalApi +import okhttp3.internal.dns.DnsCallStateMachine.Transport import okhttp3.internal.dns.DnsCallStateMachineTester.Event.OnRecords import okhttp3.internal.dns.DnsCallStateMachineTester.Event.QueryEnqueued +import okhttp3.internal.dns.DnsMessage.Companion.query import okio.ByteString /** @@ -46,11 +48,14 @@ class DnsCallStateMachineTester internal constructor( private var acceptCallbacks: Boolean = true private val transport = - object : DnsCallStateMachine.Transport { - override fun newQuery(dnsMessage: DnsMessage) = Query(dnsMessage) + object : Transport { + override fun newQuery(question: Question) = Query(question) - override fun enqueue(query: Query) { - postEvent(QueryEnqueued(query)) + override fun enqueue( + query: Query, + callback: Transport.Callback + ) { + postEvent(QueryEnqueued(query, callback)) } override fun cancel(query: Query) { @@ -157,64 +162,6 @@ class DnsCallStateMachineTester internal constructor( return event } - /** 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, - ) - }, - ), - ) - } - - /** 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, - ), - ), - ), - ) - } - - /** Respond to any query with a failure. */ - fun respondFailure( - query: Query, - e: IOException, - ) { - stateMachine.onQueryFailure(query, e) - } - /** * Asserts that the next-posted event is a call to [Dns.Callback.onRecords] with a list of IP * addresses. @@ -255,46 +202,80 @@ class DnsCallStateMachineTester internal constructor( } class Query( - val dnsMessage: DnsMessage, + val question: Question, ) sealed interface Event { - data class QueryEnqueued( + class QueryEnqueued( val query: Query, + val callback: Transport.Callback, ) : Event { 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 { 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( + class OnRecords( val last: Boolean, val records: List, ) : Event - data class OnFailure( + class OnFailure( val e: IOException, ) : Event } From 96ac1dc59391f0d6e874eaec904dea688dfcb7fc Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Wed, 22 Jul 2026 18:14:07 -0400 Subject: [PATCH 2/5] New CachingTransport --- .../okhttp3/internal/concurrent/TaskFaker.kt | 10 + .../okhttp3/internal/dns/CachingTransport.kt | 263 ++++++++++++++ .../internal/dns/DnsCallStateMachineTest.kt | 322 +++++++++++++---- .../internal/dns/DnsCallStateMachineTester.kt | 340 ++++++++++-------- 4 files changed, 722 insertions(+), 213 deletions(-) create mode 100644 okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt 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..52dd15b97f51 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,10 @@ import java.util.concurrent.BlockingQueue import java.util.concurrent.Executors import java.util.concurrent.TimeUnit import java.util.logging.Logger +import kotlin.time.Clock +import kotlin.time.Duration.Companion.nanoseconds +import kotlin.time.ExperimentalTime +import kotlin.time.Instant import okhttp3.TestUtil.threadFactory /** @@ -82,6 +87,11 @@ class TaskFaker : Closeable { /** Guarded by `this`. */ private var activeThreads = 0 + /** Adapt this API to Kotlin's time API. */ + val clock = object : Clock { + override fun now(): Instant = Instant.fromEpochSeconds(0L) + nanoTime.nanoseconds + } + /** 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/CachingTransport.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt new file mode 100644 index 000000000000..aa9b4e12300e --- /dev/null +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt @@ -0,0 +1,263 @@ +/* + * 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.Clock +import kotlin.time.Duration +import kotlin.time.Duration.Companion.seconds +import kotlin.time.ExperimentalTime +import kotlin.time.Instant +import okhttp3.internal.OkHttpInternalApi +import okhttp3.internal.concurrent.TaskRunner +import okhttp3.internal.dns.DnsCallStateMachine.Transport + +/** + * 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. + */ +// TODO: evict old entries from cache using State.lastRequestedAt +@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 clock: Clock, + 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 entry = entries.computeIfAbsent(question) { Entry(question) } + 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 = clock.now() + 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: Instant? = null, + val inFlightCall: InFlightCall? = null, + val result: Result? = null, + ) + + /** A call to the underlying transport. */ + data class InFlightCall( + val query: Q, + val sentAt: Instant, + /** The possibly-empty set of queries to notify when this call is complete. */ + val queries: List>, + ) + + /** A cached result. */ + sealed interface Result { + val revalidateAt: Instant + val expireAt: Instant + + class Failure( + override val revalidateAt: Instant, + override val expireAt: Instant, + val exception: IOException, + ) : Result + + class Success( + override val revalidateAt: Instant, + override val expireAt: Instant, + 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 bf3ce4371f66..b8a67e60cc3c 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt @@ -17,75 +17,207 @@ package okhttp3.internal.dns +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 { + private val ipv6_1_2_3_4 = listOf(InetAddress.getByName("1:2::3:4")) + private val ipv4_10_20_30_40 = 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) query1.respondIpAddresses( - addresses = listOf(InetAddress.getByName("1:2::3:4")), + addresses = ipv6_1_2_3_4, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("1:2::3:4")), + call.takeOnRecordsIpAddresses( + addresses = ipv6_1_2_3_4, ) query2.respondIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + addresses = ipv4_10_20_30_40, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + call.takeOnRecordsIpAddresses( + addresses = ipv4_10_20_30_40, ) query0.respondServiceMetadata( alpnIds = listOf("h2"), ) - takeOnRecordsServiceMetadata( + call.takeOnRecordsServiceMetadata( last = true, alpnIds = listOf(Protocol.HTTP_2), ) } + } @Test - fun `failure returned last`() = - testDnsCallStateMachine( + fun `cache already completed values`() = testDnsCallStateMachine { + val call0 = newCall( request = Dns.Request(hostname = "lysine.dev"), - ) { - enqueue() + 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 = ipv6_1_2_3_4, + ) + call0.takeOnRecordsIpAddresses( + addresses = ipv6_1_2_3_4, + ) + + call0QueryIpv4.respondIpAddresses( + addresses = ipv4_10_20_30_40, + ) + call0.takeOnRecordsIpAddresses( + last = true, + addresses = ipv4_10_20_30_40, + ) + + val call1 = newCall( + request = Dns.Request(hostname = "lysine.dev"), + includeServiceMetadata = false, + caching = true, + ) + call1.enqueue() + + call1.takeOnRecordsIpAddresses( + addresses = ipv6_1_2_3_4, + ) + call1.takeOnRecordsIpAddresses( + last = true, + addresses = ipv4_10_20_30_40, + ) + } + + /** Confirm that two queries to the cache yield a single query to the underlying transport. */ + @Test + fun `cache in flight calls`() = 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) + + val call1 = newCall( + request = Dns.Request(hostname = "lysine.dev"), + includeServiceMetadata = false, + caching = true, + ) + call1.enqueue() + + call0QueryIpv6.respondIpAddresses( + addresses = ipv6_1_2_3_4, + ) + call0.takeOnRecordsIpAddresses( + addresses = ipv6_1_2_3_4, + ) + call1.takeOnRecordsIpAddresses( + addresses = ipv6_1_2_3_4, + ) + + call0QueryIpv4.respondIpAddresses( + addresses = ipv4_10_20_30_40, + ) + call0.takeOnRecordsIpAddresses( + last = true, + addresses = ipv4_10_20_30_40, + ) + call1.takeOnRecordsIpAddresses( + last = true, + addresses = ipv4_10_20_30_40, + ) + } - val query0 = takeQuery("lysine.dev", TYPE_HTTPS) - val query1 = takeQuery("lysine.dev", TYPE_AAAA) - val query2 = takeQuery("lysine.dev", TYPE_A) + @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 = listOf(InetAddress.getByName("10.20.30.40")), + addresses = ipv4_10_20_30_40, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + call.takeOnRecordsIpAddresses( + addresses = ipv4_10_20_30_40, ) 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 = ipv4_10_20_30_40, + ) + + call0.takeOnRecordsIpAddresses( + addresses = ipv4_10_20_30_40, + ) + call0.takeOnFailure("boom!") + + val call1 = newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + includeServiceMetadata = false, + ) + call1.enqueue() + + call1.takeOnRecordsIpAddresses( + addresses = ipv4_10_20_30_40, + ) + call1.takeOnFailure("boom!") } /** @@ -96,29 +228,33 @@ 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 = { query2.respondIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + addresses = ipv4_10_20_30_40, ) query1.respondIpAddresses( - addresses = listOf(InetAddress.getByName("1:2::3:4")), + addresses = ipv6_1_2_3_4, ) } query0.respondServiceMetadata( alpnIds = listOf("h2"), ) - takeOnRecordsServiceMetadata( + call.takeOnRecordsServiceMetadata( alpnIds = listOf(Protocol.HTTP_2), ) - takeOnRecordsIpAddresses( + call.takeOnRecordsIpAddresses( last = true, addresses = listOf( @@ -134,63 +270,121 @@ 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() - enqueue() + call.enqueue() - takeCancel("lysine.dev", TYPE_HTTPS) - val query0 = takeQuery("lysine.dev", TYPE_HTTPS) - takeCancel("lysine.dev", TYPE_AAAA) - val query1 = takeQuery("lysine.dev", TYPE_AAAA) - takeCancel("lysine.dev", TYPE_A) - val query2 = takeQuery("lysine.dev", TYPE_A) + 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") - takeOnFailure("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() + 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) + + query0.respondFailure("canceled") + query1.respondFailure("canceled") + query2.respondFailure("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) query1.respondIpAddresses( - addresses = listOf(InetAddress.getByName("1:2::3:4")), + addresses = ipv6_1_2_3_4, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("1:2::3:4")), + call.takeOnRecordsIpAddresses( + addresses = ipv6_1_2_3_4, ) call.cancel() - takeCancel("lysine.dev", TYPE_HTTPS) - takeCancel("lysine.dev", TYPE_A) + transport.takeCancel("lysine.dev", TYPE_HTTPS) + transport.takeCancel("lysine.dev", TYPE_A) query2.respondIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + addresses = ipv4_10_20_30_40, ) - takeOnRecordsIpAddresses( - addresses = listOf(InetAddress.getByName("10.20.30.40")), + call.takeOnRecordsIpAddresses( + addresses = ipv4_10_20_30_40, ) 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 = ipv6_1_2_3_4, + ) + call.takeOnRecordsIpAddresses( + addresses = ipv6_1_2_3_4, + ) + + call.cancel() + + query2.respondIpAddresses( + addresses = ipv4_10_20_30_40, + ) + + 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 176c45ed32db..034f826dd04d 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt @@ -1,20 +1,25 @@ -@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.concurrent.TaskFaker import okhttp3.internal.dns.DnsCallStateMachine.Transport -import okhttp3.internal.dns.DnsCallStateMachineTester.Event.OnRecords -import okhttp3.internal.dns.DnsCallStateMachineTester.Event.QueryEnqueued +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 @@ -27,189 +32,224 @@ import okio.ByteString * calling callbacks. */ fun testDnsCallStateMachine( - request: Dns.Request, - includeIPv6: Boolean = true, - includeServiceMetadata: Boolean = true, block: DnsCallStateMachineTester.() -> Unit, ) { - val tester = DnsCallStateMachineTester(request, includeIPv6, includeServiceMetadata) + 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 : Transport { - override fun newQuery(question: Question) = Query(question) + val transport = Transport() - override fun enqueue( - query: Query, - callback: Transport.Callback - ) { - postEvent(QueryEnqueued(query, callback)) - } + private val taskFaker = TaskFaker() - override fun cancel(query: Query) { - postEvent(Event.QueryCanceled(query)) - } + private val cachingTransport = CachingTransport( + taskRunner = taskFaker.taskRunner, + delegate = transport, + clock = taskFaker.clock, + ) + + fun newCall( + request: Dns.Request, + includeIPv6: Boolean = true, + includeServiceMetadata: Boolean = true, + caching: Boolean = false, + ): Call = Call( + request = request, + includeIPv6 = includeIPv6, + includeServiceMetadata = includeServiceMetadata, + caching = caching + ) + + inner class Transport : DnsCallStateMachine.Transport { + val events = LinkedBlockingDeque() + + override fun newQuery(question: Question) = Query(question) + + private fun postEvent(e: TransportEvent) { + events.put(e) + + onNextEvent?.invoke() + onNextEvent = null } - val call: Dns.Call = - object : Dns.Call { - override val request: Dns.Request = request + private fun takeEvent(): TransportEvent { + taskFaker.runTasks() // Run any queued async work first. + return events.take() + } - 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 = transport.takeEvent() as QueryEnqueued + assertThat(event.hostname).isEqualTo(hostname) + assertThat(event.type).isEqualTo(type) + return event + } + /** 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 + } - override fun cancel() { - stateMachine.cancel() - } + override fun enqueue( + query: Query, + callback: Transport.Callback + ) { + postEvent(QueryEnqueued(query, callback)) + } - override fun isCanceled() = stateMachine.canceled + override fun cancel(query: Query) { + postEvent(QueryCanceled(query)) } + } - 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 - } + /** 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 } + } - 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 - } - } + override fun enqueue(callback: Dns.Callback) { + stateMachine.start(callback) } - val stateMachine = - DnsCallStateMachine( - transport = transport, - call = call, - canceledException = null, - includeIPv6 = includeIPv6, - includeServiceMetadata = includeServiceMetadata, - ) - - /** Start the DNS call. */ - fun enqueue() { - check(acceptCallbacks) { "unexpected enqueue" } - - acceptCallbacks = false - try { - call.enqueue(callback) - } finally { - acceptCallbacks = true + override fun cancel() { + stateMachine.cancel() } - } - private fun postEvent(e: Event) { - events.put(e) + override fun isCanceled() = stateMachine.canceled - onNextEvent?.invoke() - onNextEvent = null - } + private fun postEvent(e: CallEvent) { + events.put(e) - /** 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 - } + onNextEvent?.invoke() + onNextEvent = null + } - /** 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 - } + 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.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): 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.onFailure]. */ + fun takeOnFailure(message: String): OnFailure { + val event = takeEvent() as OnFailure + assertThat(event.e).hasMessage(message) + return event + } } class Query( val question: Question, ) - sealed interface Event { + sealed interface TransportEvent { class QueryEnqueued( val query: Query, val callback: Transport.Callback, - ) : Event { + ) : TransportEvent { val hostname: String get() = query.question.name val type: Int @@ -263,20 +303,22 @@ class DnsCallStateMachineTester internal constructor( class QueryCanceled( val query: Query, - ) : Event { + ) : TransportEvent { val hostname: String get() = query.question.name val type: Int get() = query.question.type } + } + sealed interface CallEvent { class OnRecords( val last: Boolean, val records: List, - ) : Event + ) : CallEvent class OnFailure( val e: IOException, - ) : Event + ) : CallEvent } } From fed1cd6dbecf5327daa2faf438043df0063853ff Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Wed, 22 Jul 2026 19:08:12 -0400 Subject: [PATCH 3/5] spotless --- .../internal/-DnsOverHttpsCall.kt | 14 ++++-- .../okhttp3/internal/concurrent/TaskFaker.kt | 7 +-- .../internal/dns/-DnsCallStateMachine.kt | 37 +++++++++------- .../internal/dns/DnsCallStateMachineTester.kt | 43 ++++++++++--------- 4 files changed, 58 insertions(+), 43 deletions(-) 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 a293676ae54c..c2924423001a 100644 --- a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsCall.kt +++ b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsCall.kt @@ -98,15 +98,21 @@ internal class DnsOverHttpsCall( override fun enqueue( query: Call, - callback: DnsCallStateMachine.Transport.Callback + callback: DnsCallStateMachine.Transport.Callback, ) { query.enqueue( object : Callback { - override fun onFailure(call: Call, e: IOException) { + override fun onFailure( + call: Call, + e: IOException, + ) { callback.onFailure(e) } - override fun onResponse(call: Call, response: Response) { + override fun onResponse( + call: Call, + response: Response, + ) { val dnsMessage = try { decodeResponse(response) @@ -116,7 +122,7 @@ internal class DnsOverHttpsCall( callback.onResponse(dnsMessage) } - } + }, ) } 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 52dd15b97f51..62d4d9a0db2a 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 @@ -88,9 +88,10 @@ class TaskFaker : Closeable { private var activeThreads = 0 /** Adapt this API to Kotlin's time API. */ - val clock = object : Clock { - override fun now(): Instant = Instant.fromEpochSeconds(0L) + nanoTime.nanoseconds - } + val clock = + object : Clock { + override fun now(): Instant = Instant.fromEpochSeconds(0L) + nanoTime.nanoseconds + } /** A task runner that posts tasks to this fake. Tasks won't be executed until requested. */ val 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 e6ee3fe30148..267f73585fe1 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsCallStateMachine.kt @@ -103,21 +103,22 @@ class DnsCallStateMachine( 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, - ) - } - } + callback = + object : Transport.Callback { + override fun onResponse(dnsResponse: DnsMessage) { + updateStateAndCallCallbacks( + completedQuery = query, + dnsResponse = dnsResponse, + ) + } + + override fun onFailure(e: IOException) { + updateStateAndCallCallbacks( + completedQuery = query, + newException = e, + ) + } + }, ) } @@ -332,12 +333,16 @@ class DnsCallStateMachine( interface Transport { fun newQuery(question: Question): Q - fun enqueue(query: Q, callback: Callback) + 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/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt index 034f826dd04d..b4c26439ce19 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt @@ -31,9 +31,7 @@ import okio.ByteString * This tracks all effects from the state machine as events: creating queries, canceling queries, * calling callbacks. */ -fun testDnsCallStateMachine( - block: DnsCallStateMachineTester.() -> Unit, -) { +fun testDnsCallStateMachine(block: DnsCallStateMachineTester.() -> Unit) { val tester = DnsCallStateMachineTester() tester.block() assertThat(tester.transport.events.poll(), "unexpected transport event").isNull() @@ -49,23 +47,25 @@ class DnsCallStateMachineTester internal constructor() { private val taskFaker = TaskFaker() - private val cachingTransport = CachingTransport( - taskRunner = taskFaker.taskRunner, - delegate = transport, - clock = taskFaker.clock, - ) + private val cachingTransport = + CachingTransport( + taskRunner = taskFaker.taskRunner, + delegate = transport, + clock = taskFaker.clock, + ) fun newCall( request: Dns.Request, includeIPv6: Boolean = true, includeServiceMetadata: Boolean = true, caching: Boolean = false, - ): Call = Call( - request = request, - includeIPv6 = includeIPv6, - includeServiceMetadata = includeServiceMetadata, - caching = caching - ) + ): Call = + Call( + request = request, + includeIPv6 = includeIPv6, + includeServiceMetadata = includeServiceMetadata, + caching = caching, + ) inner class Transport : DnsCallStateMachine.Transport { val events = LinkedBlockingDeque() @@ -94,6 +94,7 @@ class DnsCallStateMachineTester internal constructor() { assertThat(event.type).isEqualTo(type) return event } + /** Asserts that the next-posted event is a query cancel. */ fun takeCancel( hostname: String, @@ -107,7 +108,7 @@ class DnsCallStateMachineTester internal constructor() { override fun enqueue( query: Query, - callback: Transport.Callback + callback: Transport.Callback, ) { postEvent(QueryEnqueued(query, callback)) } @@ -123,15 +124,17 @@ class DnsCallStateMachineTester internal constructor() { includeIPv6: Boolean = true, includeServiceMetadata: Boolean = true, caching: Boolean = false, - ) : Dns.Call, Dns.Callback { + ) : Dns.Call, + Dns.Callback { private val events = LinkedBlockingDeque() val stateMachine = DnsCallStateMachine( - transport = when { - caching -> cachingTransport - else -> transport - }, + transport = + when { + caching -> cachingTransport + else -> transport + }, call = this, canceledException = null, includeIPv6 = includeIPv6, From 99a441ab7de708368816753490600363d0bed95e Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Thu, 23 Jul 2026 09:46:25 -0400 Subject: [PATCH 4/5] Switch to ComparableTimeSource --- .../okhttp3/internal/concurrent/TaskFaker.kt | 11 +- .../okhttp3/internal/dns/CachingTransport.kt | 117 ++++--- .../internal/dns/DnsCallStateMachineTest.kt | 294 +++++++++--------- .../internal/dns/DnsCallStateMachineTester.kt | 2 +- 4 files changed, 230 insertions(+), 194 deletions(-) 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 62d4d9a0db2a..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 @@ -31,10 +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.Clock -import kotlin.time.Duration.Companion.nanoseconds +import kotlin.time.AbstractLongTimeSource +import kotlin.time.DurationUnit import kotlin.time.ExperimentalTime -import kotlin.time.Instant import okhttp3.TestUtil.threadFactory /** @@ -88,9 +87,9 @@ class TaskFaker : Closeable { private var activeThreads = 0 /** Adapt this API to Kotlin's time API. */ - val clock = - object : Clock { - override fun now(): Instant = Instant.fromEpochSeconds(0L) + nanoTime.nanoseconds + 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. */ diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt index aa9b4e12300e..58351b8de2fe 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt @@ -18,15 +18,17 @@ package okhttp3.internal.dns import java.io.IOException import java.util.concurrent.ConcurrentHashMap import java.util.concurrent.atomic.AtomicReference -import kotlin.time.Clock +import kotlin.time.ComparableTimeMark as Time import kotlin.time.Duration import kotlin.time.Duration.Companion.seconds import kotlin.time.ExperimentalTime -import kotlin.time.Instant +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. @@ -48,13 +50,12 @@ import okhttp3.internal.dns.DnsCallStateMachine.Transport * If this receives multiple equivalent queries, it combines them into a single query on the * underlying transport. */ -// TODO: evict old entries from cache using State.lastRequestedAt @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 clock: Clock, + private val timeSource: TimeSource.WithComparableMarks, private val minimumTimeToLive: Duration = 10.seconds, private val maximumTimeToLive: Duration = 300.seconds, private val failureTimeToLive: Duration = 10.seconds, @@ -76,13 +77,13 @@ class CachingTransport( override fun enqueue( query: Query, - callback: Transport.Callback> + callback: Transport.Callback>, ) { check(query.callback == null) { "already enqueued" } query.callback = callback val entry = query.entry - val now = clock.now() + val now = timeSource.markNow() while (true) { val previous = entry.state.get() val result = previous.result @@ -93,20 +94,32 @@ class CachingTransport( // 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 - }, - ) + 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. @@ -141,11 +154,13 @@ class CachingTransport( val newQueries = inFlightCall.queries - query if (newQueries.size == inFlightCall.queries.size) return - val next = previous.copy( - inFlightCall = inFlightCall.copy( - queries = newQueries, + val next = + previous.copy( + inFlightCall = + inFlightCall.copy( + queries = newQueries, + ), ) - ) if (!entry.state.compareAndSet(previous, next)) continue // Lost a race, retry. @@ -179,14 +194,16 @@ class CachingTransport( 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, + 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. @@ -203,18 +220,22 @@ class CachingTransport( while (true) { val previous = state.get() val sentAt = previous.inFlightCall!!.sentAt - val timeToLive = (dnsResponse.answers.minOfOrNull { it.timeToLive } ?: 0).seconds - .coerceIn(minimumTimeToLive, maximumTimeToLive) + 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, + 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. @@ -230,7 +251,7 @@ class CachingTransport( /** A snapshot of the state of a single entry. */ data class State( - val lastRequestedAt: Instant? = null, + val lastRequestedAt: Time? = null, val inFlightCall: InFlightCall? = null, val result: Result? = null, ) @@ -238,25 +259,25 @@ class CachingTransport( /** A call to the underlying transport. */ data class InFlightCall( val query: Q, - val sentAt: Instant, + 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: Instant - val expireAt: Instant + val revalidateAt: Time + val expireAt: Time class Failure( - override val revalidateAt: Instant, - override val expireAt: Instant, + override val revalidateAt: Time, + override val expireAt: Time, val exception: IOException, ) : Result class Success( - override val revalidateAt: Instant, - override val expireAt: Instant, + 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 b8a67e60cc3c..6a5b162afd81 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTest.kt @@ -26,16 +26,18 @@ import okhttp3.internal.OkHttpInternalApi @Burst class DnsCallStateMachineTest { - private val ipv6_1_2_3_4 = listOf(InetAddress.getByName("1:2::3:4")) - private val ipv4_10_20_30_40 = listOf(InetAddress.getByName("10.20.30.40")) + /** 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`(caching: Boolean = true) { testDnsCallStateMachine { - val call = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = caching, - ) + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = caching, + ) call.enqueue() val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) @@ -43,17 +45,17 @@ class DnsCallStateMachineTest { val query2 = transport.takeQuery("lysine.dev", TYPE_A) query1.respondIpAddresses( - addresses = ipv6_1_2_3_4, + addresses = blueIpv6s, ) call.takeOnRecordsIpAddresses( - addresses = ipv6_1_2_3_4, + addresses = blueIpv6s, ) query2.respondIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) call.takeOnRecordsIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) query0.respondServiceMetadata( @@ -67,98 +69,105 @@ class DnsCallStateMachineTest { } @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 = ipv6_1_2_3_4, - ) - call0.takeOnRecordsIpAddresses( - addresses = ipv6_1_2_3_4, - ) - - call0QueryIpv4.respondIpAddresses( - addresses = ipv4_10_20_30_40, - ) - call0.takeOnRecordsIpAddresses( - last = true, - addresses = ipv4_10_20_30_40, - ) - - val call1 = newCall( - request = Dns.Request(hostname = "lysine.dev"), - includeServiceMetadata = false, - caching = true, - ) - call1.enqueue() - - call1.takeOnRecordsIpAddresses( - addresses = ipv6_1_2_3_4, - ) - call1.takeOnRecordsIpAddresses( - last = true, - addresses = ipv4_10_20_30_40, - ) - } + 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 `cache in flight calls`() = 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) - - val call1 = newCall( - request = Dns.Request(hostname = "lysine.dev"), - includeServiceMetadata = false, - caching = true, - ) - call1.enqueue() - - call0QueryIpv6.respondIpAddresses( - addresses = ipv6_1_2_3_4, - ) - call0.takeOnRecordsIpAddresses( - addresses = ipv6_1_2_3_4, - ) - call1.takeOnRecordsIpAddresses( - addresses = ipv6_1_2_3_4, - ) - - call0QueryIpv4.respondIpAddresses( - addresses = ipv4_10_20_30_40, - ) - call0.takeOnRecordsIpAddresses( - last = true, - addresses = ipv4_10_20_30_40, - ) - call1.takeOnRecordsIpAddresses( - last = true, - addresses = ipv4_10_20_30_40, - ) - } + fun `cache in flight calls`() = + 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) + + 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, + ) + + call0QueryIpv4.respondIpAddresses( + addresses = blueIpv4s, + ) + call0.takeOnRecordsIpAddresses( + last = true, + addresses = blueIpv4s, + ) + call1.takeOnRecordsIpAddresses( + last = true, + addresses = blueIpv4s, + ) + } @Test fun `failure returned last`(caching: Boolean = true) = testDnsCallStateMachine { - val call = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = caching, - ) + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = caching, + ) call.enqueue() val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) @@ -168,10 +177,10 @@ class DnsCallStateMachineTest { query1.respondFailure("boom!") query2.respondIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) call.takeOnRecordsIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) query0.respondServiceMetadata( @@ -187,11 +196,12 @@ class DnsCallStateMachineTest { @Test fun `failure is cached`() = testDnsCallStateMachine { - val call0 = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = true, - includeServiceMetadata = false, - ) + val call0 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + includeServiceMetadata = false, + ) call0.enqueue() val queryIpv6 = transport.takeQuery("lysine.dev", TYPE_AAAA) @@ -199,23 +209,24 @@ class DnsCallStateMachineTest { queryIpv6.respondFailure("boom!") queryIpv4.respondIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) call0.takeOnRecordsIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) call0.takeOnFailure("boom!") - val call1 = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = true, - includeServiceMetadata = false, - ) + val call1 = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + includeServiceMetadata = false, + ) call1.enqueue() call1.takeOnRecordsIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) call1.takeOnFailure("boom!") } @@ -230,10 +241,11 @@ class DnsCallStateMachineTest { @Test fun `calls to onRecords are serialized`(caching: Boolean = true) = testDnsCallStateMachine { - val call = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = caching, - ) + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = caching, + ) call.enqueue() val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) @@ -242,10 +254,10 @@ class DnsCallStateMachineTest { onNextEvent = { query2.respondIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) query1.respondIpAddresses( - addresses = ipv6_1_2_3_4, + addresses = blueIpv6s, ) } query0.respondServiceMetadata( @@ -271,10 +283,11 @@ class DnsCallStateMachineTest { @Test fun `cancel before enqueue`() = testDnsCallStateMachine { - val call = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = false, - ) + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = false, + ) call.cancel() call.enqueue() @@ -295,10 +308,11 @@ class DnsCallStateMachineTest { @Test fun `cancel before enqueue with caching`() = testDnsCallStateMachine { - val call = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = true, - ) + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + ) call.cancel() call.enqueue() @@ -317,10 +331,11 @@ class DnsCallStateMachineTest { @Test fun `cancel ignored if canceled query completes`() = testDnsCallStateMachine { - val call = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = false, - ) + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = false, + ) call.enqueue() val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) @@ -328,10 +343,10 @@ class DnsCallStateMachineTest { val query2 = transport.takeQuery("lysine.dev", TYPE_A) query1.respondIpAddresses( - addresses = ipv6_1_2_3_4, + addresses = blueIpv6s, ) call.takeOnRecordsIpAddresses( - addresses = ipv6_1_2_3_4, + addresses = blueIpv6s, ) call.cancel() @@ -340,10 +355,10 @@ class DnsCallStateMachineTest { transport.takeCancel("lysine.dev", TYPE_A) query2.respondIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) call.takeOnRecordsIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) query0.respondServiceMetadata( @@ -359,10 +374,11 @@ class DnsCallStateMachineTest { @Test fun `cancel ignored if canceled query completes with caching`() = testDnsCallStateMachine { - val call = newCall( - request = Dns.Request(hostname = "lysine.dev"), - caching = true, - ) + val call = + newCall( + request = Dns.Request(hostname = "lysine.dev"), + caching = true, + ) call.enqueue() val query0 = transport.takeQuery("lysine.dev", TYPE_HTTPS) @@ -370,16 +386,16 @@ class DnsCallStateMachineTest { val query2 = transport.takeQuery("lysine.dev", TYPE_A) query1.respondIpAddresses( - addresses = ipv6_1_2_3_4, + addresses = blueIpv6s, ) call.takeOnRecordsIpAddresses( - addresses = ipv6_1_2_3_4, + addresses = blueIpv6s, ) call.cancel() query2.respondIpAddresses( - addresses = ipv4_10_20_30_40, + addresses = blueIpv4s, ) query0.respondServiceMetadata( diff --git a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt index b4c26439ce19..0d122444c4a7 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsCallStateMachineTester.kt @@ -51,7 +51,7 @@ class DnsCallStateMachineTester internal constructor() { CachingTransport( taskRunner = taskFaker.taskRunner, delegate = transport, - clock = taskFaker.clock, + timeSource = taskFaker.timeSource, ) fun newCall( From d954f054a0bd9d3d69875b983b58e32c674714b3 Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Thu, 23 Jul 2026 10:00:15 -0400 Subject: [PATCH 5/5] Switch from computeIfAbsent to putIfAbsent The former isn't available on API 21. --- .../kotlin/okhttp3/internal/dns/CachingTransport.kt | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt index 58351b8de2fe..1ad0c04f3d5a 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/CachingTransport.kt @@ -71,7 +71,8 @@ class CachingTransport( } override fun newQuery(question: Question): Query { - val entry = entries.computeIfAbsent(question) { Entry(question) } + val inserted = Entry(question) + val entry = entries.putIfAbsent(question, inserted) ?: inserted return Query(entry) }