diff --git a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/DnsOverHttps.kt b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/DnsOverHttps.kt index ed543bf18c55..6fe03cefcbd5 100644 --- a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/DnsOverHttps.kt +++ b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/DnsOverHttps.kt @@ -22,7 +22,7 @@ import okhttp3.HttpUrl import okhttp3.MediaType import okhttp3.MediaType.Companion.toMediaType import okhttp3.OkHttpClient -import okhttp3.dnsoverhttps.internal.DnsOverHttpsTransport +import okhttp3.dnsoverhttps.internal.DnsOverHttpsQuery import okhttp3.internal.dns.StateMachineDnsCall import okhttp3.internal.dns.execute import okhttp3.internal.publicsuffix.PublicSuffixDatabase @@ -46,8 +46,8 @@ class DnsOverHttps internal constructor( @get:JvmName("resolvePrivateAddresses") val resolvePrivateAddresses: Boolean, @get:JvmName("resolvePublicAddresses") val resolvePublicAddresses: Boolean, ) : Dns { - private val transport = - DnsOverHttpsTransport( + private val queryFactory = + DnsOverHttpsQuery.Factory( client = client, dnsUrl = url, post = post, @@ -56,7 +56,7 @@ class DnsOverHttps internal constructor( override fun newCall(request: Dns.Request): Dns.Call = StateMachineDnsCall( request = request, - transport = transport, + queryFactory = queryFactory, canceledException = validate(request.hostname), includeIPv6 = includeIPv6, includeServiceMetadata = includeServiceMetadata, diff --git a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/DnsOverHttpsTransport.kt b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsQuery.kt similarity index 66% rename from okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/DnsOverHttpsTransport.kt rename to okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsQuery.kt index e735af3b4a30..4e721eecef25 100644 --- a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/DnsOverHttpsTransport.kt +++ b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsOverHttpsQuery.kt @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ +@file:Suppress("ktlint:standard:filename") + package okhttp3.dnsoverhttps.internal import java.io.IOException @@ -25,60 +27,24 @@ import okhttp3.Protocol import okhttp3.Request import okhttp3.RequestBody import okhttp3.Response -import okhttp3.dnsoverhttps.DnsOverHttps import okhttp3.dnsoverhttps.DnsOverHttps.Companion.DNS_MESSAGE import okhttp3.dnsoverhttps.DnsOverHttps.Companion.MAX_RESPONSE_SIZE import okhttp3.internal.OkHttpInternalApi import okhttp3.internal.dns.DnsMessage import okhttp3.internal.dns.DnsMessageReader import okhttp3.internal.dns.DnsMessageWriter +import okhttp3.internal.dns.DnsQuery import okhttp3.internal.dns.Question -import okhttp3.internal.dns.StateMachineDnsCall import okhttp3.internal.platform.Platform import okio.Buffer import okio.BufferedSink @OkHttpInternalApi -internal class DnsOverHttpsTransport( - private val client: OkHttpClient, - private val dnsUrl: HttpUrl, - private val post: Boolean, -) : StateMachineDnsCall.Transport { - override fun newQuery(question: Question): Call { - val dnsMessage = DnsMessage.query(question) - return client.newCall( - request = - Request - .Builder() - .header("Accept", DnsOverHttps.DNS_MESSAGE.toString()) - .apply { - if (post) { - url(dnsUrl) - cacheUrlOverride( - dnsUrl - .newBuilder() - .addQueryParameter("hostname", question.name) - .addQueryParameter("type", question.type.toString()) - .build(), - ) - post(QueryRequestBody(dnsMessage)) - } else { - val requestUrl = - dnsUrl - .newBuilder() - .addQueryParameter("dns", dnsMessage.asQueryParameter()) - .build() - url(requestUrl) - } - }.build(), - ) - } - - override fun enqueue( - query: Call, - callback: StateMachineDnsCall.Transport.Callback, - ) { - query.enqueue( +internal class DnsOverHttpsQuery( + val call: Call, +) : DnsQuery { + override fun enqueue(callback: DnsQuery.Callback) { + call.enqueue( object : Callback { override fun onFailure( call: Call, @@ -104,8 +70,47 @@ internal class DnsOverHttpsTransport( ) } - override fun cancel(query: Call) { - query.cancel() + override fun cancel() { + call.cancel() + } + + class Factory( + private val client: OkHttpClient, + private val dnsUrl: HttpUrl, + private val post: Boolean, + ) : DnsQuery.Factory { + override fun newQuery(question: Question): DnsQuery { + val dnsMessage = DnsMessage.query(question) + return DnsOverHttpsQuery( + call = + client.newCall( + request = + Request + .Builder() + .header("Accept", DNS_MESSAGE.toString()) + .apply { + if (post) { + url(dnsUrl) + cacheUrlOverride( + dnsUrl + .newBuilder() + .addQueryParameter("hostname", question.name) + .addQueryParameter("type", question.type.toString()) + .build(), + ) + post(QueryRequestBody(dnsMessage)) + } else { + val requestUrl = + dnsUrl + .newBuilder() + .addQueryParameter("dns", dnsMessage.asQueryParameter()) + .build() + url(requestUrl) + } + }.build(), + ), + ) + } } } diff --git a/okhttp-testing-support/src/main/kotlin/okhttp3/FakeDns.kt b/okhttp-testing-support/src/main/kotlin/okhttp3/FakeDns.kt index f81392d16689..78fff4e85481 100644 --- a/okhttp-testing-support/src/main/kotlin/okhttp3/FakeDns.kt +++ b/okhttp-testing-support/src/main/kotlin/okhttp3/FakeDns.kt @@ -195,7 +195,11 @@ class FakeDns( } } - return dnsResponse(request, answers) + return DnsMessage.response( + id = request.id, + questions = request.questions, + answers = answers, + ) } private fun ResourceRecord.matches(question: Question): Boolean { @@ -323,23 +327,3 @@ class FakeDns( ) : Request } } - -fun dnsResponse( - request: DnsMessage, - answers: List, -): DnsMessage { - // QR = 1 (Response) - // OPCODE = 0 (standard query) - // RD = 1 (Recursion Desired) - // RA = 1 (Recursion Available) - // RCODE = 0 (success) - // QR OPCODE AA TC RD RA Z RCODE - val flags = 0b1___0000__0__0__1__1_000__0000 - - return DnsMessage( - id = request.id, - flags = flags, - questions = request.questions, - answers = answers, - ) -} diff --git a/okhttp/src/androidMain/kotlin/okhttp3/android/AndroidDns.kt b/okhttp/src/androidMain/kotlin/okhttp3/android/AndroidDns.kt index aabf96cc8ba2..65ac4b304831 100644 --- a/okhttp/src/androidMain/kotlin/okhttp3/android/AndroidDns.kt +++ b/okhttp/src/androidMain/kotlin/okhttp3/android/AndroidDns.kt @@ -29,10 +29,10 @@ import java.util.concurrent.Executor import okhttp3.Dns import okhttp3.internal.OkHttpInternalApi import okhttp3.internal.SuppressSignatureCheck -import okhttp3.internal.concurrent.Task import okhttp3.internal.concurrent.TaskRunner import okhttp3.internal.dns.DnsMessage import okhttp3.internal.dns.DnsMessageReader +import okhttp3.internal.dns.DnsQuery import okhttp3.internal.dns.Question import okhttp3.internal.dns.ResourceRecord import okhttp3.internal.dns.StateMachineDnsCall @@ -42,192 +42,177 @@ import okhttp3.internal.dns.execute import okio.Buffer /** - * A [Dns] backed by Android's system resolver, with ECH support from Android's [DnsResolver]. + * A [Dns] backed by Android's system resolver, with Encrypted Client Hello (ECH) support from + * Android's [DnsResolver]. * - * IP addresses come from the system resolver — [InetAddress.getAllByName], or - * [Network.getAllByName] when a [network] is set. Internally this drives OkHttp's - * [StateMachineDnsCall] through a private transport: a single `A` query stands in for the blocking - * address lookup (which returns both address families), plus an optional `HTTPS` query for service - * metadata such as ECH. + * IP addresses come from the system resolver — [InetAddress.getAllByName] or [Network.getAllByName] + * when [network] is non-null. It makes an optional `HTTPS` query for service metadata such as ECH. */ @RequiresApi(29) @SuppressSignatureCheck -class AndroidDns - constructor( - private val dnsResolver: DnsResolver = DnsResolver.getInstance(), - private val network: Network? = null, - /** - * True to also query the `HTTPS` record for service metadata. Keep this on: it enables privacy - * features such as Encrypted Client Hello (ECH) for the HTTPS call. Set it to false only when - * you want to disable ECH. - */ - private val includeServiceMetadata: Boolean = true, - // Runs inline; the executor only hands off DnsResolver's callbacks. - private val executor: Executor = Executor { it.run() }, - ) : Dns { - private val transport = AndroidTransport() +class AndroidDns( + private val dnsResolver: DnsResolver = DnsResolver.getInstance(), + private val network: Network? = null, + /** + * True to also query the `HTTPS` record for service metadata. Keep this on: it enables privacy + * features such as Encrypted Client Hello (ECH) for the HTTPS call. Set it to false only when + * you want to disable ECH. + */ + private val includeServiceMetadata: Boolean = true, + // Runs inline; the executor only hands off DnsResolver's callbacks. + private val executor: Executor = Executor { it.run() }, +) : Dns { + private val taskRunner: TaskRunner = TaskRunner.INSTANCE + + /** Drives [StateMachineDnsCall] using the system resolver and [DnsResolver]. */ + private val queryFactory = + DnsQuery.Factory { question -> + AndroidQuery(question) + } + + /** + * Resolves addresses only, for callers using the legacy blocking API. [Dns] cannot carry + * HTTPS/ECH metadata, so this skips that query rather than paying for it and discarding it. + */ + override fun lookup(hostname: String): List = + call(Dns.Request(hostname), includeServiceMetadata = false) + .execute() + .filterIsInstance() + .map { it.address } + + override fun newCall(request: Dns.Request): Dns.Call = call(request, includeServiceMetadata = includeServiceMetadata) + + private fun call( + request: Dns.Request, + includeServiceMetadata: Boolean, + ): Dns.Call = + StateMachineDnsCall( + request = request, + queryFactory = queryFactory, + canceledException = null, + // A single `A` query stands in for both families: the system resolver returns IPv4 and + // IPv6 addresses together, so there's no separate `AAAA` query. + includeIPv6 = false, + includeServiceMetadata = includeServiceMetadata, + ) + + /** One outstanding transport-layer query. */ + private inner class AndroidQuery( + private val question: Question, + ) : DnsQuery { + /** Only the `HTTPS` query reaches [DnsResolver], which is the API that takes a signal. */ + private val cancellationSignal = if (question.type == TYPE_HTTPS) CancellationSignal() else null + + override fun enqueue(callback: DnsQuery.Callback) { + when (question.type) { + TYPE_A -> resolveAddresses(callback) + + TYPE_HTTPS -> queryServiceMetadata(callback) + + // AndroidDns only ever issues `A` and `HTTPS` queries (includeIPv6 = false, so no `AAAA`). + else -> error("unexpected query type ${question.type}") + } + } /** - * Resolves addresses only, for callers using the legacy blocking API. [Dns] cannot carry - * HTTPS/ECH metadata, so this skips that query rather than paying for it and discarding it. + * Resolves IP addresses through the system resolver. This is a blocking call, so it runs on a + * [TaskRunner] thread. The addresses are wrapped in a synthetic [DnsMessage] because that's + * the only shape [StateMachineDnsCall] accepts — the system resolver gives us decoded + * addresses rather than a wire-format message. */ - override fun lookup(hostname: String): List = - call(Dns.Request(hostname), includeServiceMetadata = false) - .execute() - .filterIsInstance() - .map { it.address } - - override fun newCall(request: Dns.Request): Dns.Call = call(request, includeServiceMetadata = includeServiceMetadata) - - private fun call( - request: Dns.Request, - includeServiceMetadata: Boolean, - ): Dns.Call = - StateMachineDnsCall( - request = request, - transport = transport, - canceledException = null, - // A single `A` query stands in for both families: the system resolver returns IPv4 and - // IPv6 addresses together, so there's no separate `AAAA` query. - includeIPv6 = false, - includeServiceMetadata = includeServiceMetadata, - ) - - /** Drives [StateMachineDnsCall] using the system resolver and [DnsResolver]. */ - private inner class AndroidTransport : StateMachineDnsCall.Transport { - override fun newQuery(question: Question) = Query(question) - - override fun enqueue( - query: Query, - callback: StateMachineDnsCall.Transport.Callback, - ) { - when (query.type) { - TYPE_A -> resolveAddresses(query.hostname, callback) - - TYPE_HTTPS -> queryServiceMetadata(query, callback) - - // AndroidDns only ever issues `A` and `HTTPS` queries (includeIPv6 = false, so no `AAAA`). - else -> error("unexpected query type ${query.type}") + private fun resolveAddresses(callback: DnsQuery.Callback) { + taskRunner.newQueue().execute("${question.name} dns", cancelable = false) { + try { + val addresses = + network?.getAllByName(question.name) + ?: InetAddress.getAllByName(question.name) + callback.onResponse(response(addresses)) + } catch (e: UnknownHostException) { + callback.onFailure(e) } } + } - /** - * Cancels a query. Only the `HTTPS` query reaches [DnsResolver]; the address lookup runs to - * completion and its result is discarded by [StateMachineDnsCall], which ignores callbacks - * once canceled. - */ - override fun cancel(query: Query) { - query.cancellationSignal?.cancel() - } + /** + * Wraps system-resolved [addresses] in a successful [DnsMessage] so they can flow through + * [DnsQuery.Callback], which only accepts wire-format responses. + */ + private fun response(addresses: Array) = + DnsMessage.response( + questions = listOf(question), + answers = + addresses.map { address -> + ResourceRecord.IpAddress( + name = question.name, + timeToLive = 0, + address = address, + ) + }, + ) - /** - * Resolves IP addresses through the system resolver. This is a blocking call, so it runs on a - * [TaskRunner] thread. The addresses are wrapped in a synthetic [DnsMessage] because that's - * the only shape [StateMachineDnsCall] accepts — the system resolver gives us decoded - * addresses rather than a wire-format message. - */ - private fun resolveAddresses( - hostname: String, - callback: StateMachineDnsCall.Transport.Callback, - ) { - TaskRunner.INSTANCE.newQueue().schedule( - object : Task("$hostname address lookup", cancelable = false) { - override fun runOnce(): Long { + /** + * Asks [DnsResolver] for the `HTTPS` record. The platform has no typed API for this below + * API 36, so we request the raw message and decode it with OkHttp's own [DnsMessageReader] — + * the same decoder `DnsOverHttps` uses. + * + * Failures are reported to the [DnsQuery.Callback] rather than swallowed, so callers can tell + * an absent `HTTPS` record from a query that errored. + * + * `WrongConstant` is suppressed because OkHttp's [TYPE_HTTPS] is the DNS wire value (65), the + * same as the platform's `DnsResolver.TYPE_HTTPS`. We use ours so the query stays valid on + * API 29+; `DnsResolver.TYPE_HTTPS` was only added in API 37. + */ + @SuppressLint("WrongConstant") + @Suppress("ktlint:standard:comment-wrapping") + private fun queryServiceMetadata(callback: DnsQuery.Callback) { + val dnsResolverCallback = + object : DnsResolver.Callback { + override fun onAnswer( + answer: ByteArray, + rcode: Int, + ) { + val message = try { - val addresses = - when (network) { - null -> InetAddress.getAllByName(hostname) - else -> network.getAllByName(hostname) - } - callback.onResponse(addresses.toDnsMessage(hostname)) - } catch (e: UnknownHostException) { - callback.onFailure(e) + DnsMessageReader(Buffer().write(answer)).read() + } catch (e: IOException) { + return callback.onFailure(e) } - return -1L - } - }, - ) - } + // The state machine turns a non-success rcode into a failure of its own. + callback.onResponse(message) + } - /** - * Asks [DnsResolver] for the `HTTPS` record. The platform has no typed API for this below - * API 36, so we request the raw message and decode it with OkHttp's own [DnsMessageReader] — - * the same decoder `DnsOverHttps` uses. - * - * Failures are reported to the [StateMachineDnsCall] rather than swallowed, so callers can - * tell an absent `HTTPS` record from a query that errored. - * - * `WrongConstant` is suppressed because okhttp's [TYPE_HTTPS] is the DNS wire value (65), the - * same as the platform's `DnsResolver.TYPE_HTTPS`. We use ours so the query stays valid on - * API 29+; `DnsResolver.TYPE_HTTPS` was only added in API 37. - */ - @SuppressLint("WrongConstant") - @Suppress("ktlint:standard:comment-wrapping") - private fun queryServiceMetadata( - query: Query, - callback: StateMachineDnsCall.Transport.Callback, - ) { - val queryCallback = - object : DnsResolver.Callback { - override fun onAnswer( - answer: ByteArray, - rcode: Int, - ) { - val message = - try { - DnsMessageReader(Buffer().write(answer)).read() - } catch (e: IOException) { - return callback.onFailure(e) - } - // The state machine turns a non-success rcode into a failure of its own. - callback.onResponse(message) - } - - override fun onError(e: DnsResolver.DnsException) { - callback.onFailure(IOException("HTTPS query failed with code ${e.code}", e)) - } + override fun onError(e: DnsResolver.DnsException) { + callback.onFailure(IOException("HTTPS query failed with code ${e.code}", e)) } + } - try { - dnsResolver.rawQuery( - /* network = */ network, - /* domain = */ query.hostname, - /* nsClass = */ DnsResolver.CLASS_IN, - /* nsType = */ TYPE_HTTPS, - /* flags = */ DnsResolver.FLAG_EMPTY, - /* executor = */ executor, - /* cancellationSignal = */ query.cancellationSignal, - /* callback = */ queryCallback, - ) - } catch (e: SecurityException) { - // The app lacks INTERNET permission, or network policy forbids the query. The callback - // won't run, so report the failure here or the state machine never terminates. + try { + dnsResolver.rawQuery( + /* network = */ network, + /* domain = */ question.name, + /* nsClass = */ DnsResolver.CLASS_IN, + /* nsType = */ TYPE_HTTPS, + /* flags = */ DnsResolver.FLAG_EMPTY, + /* executor = */ executor, + /* cancellationSignal = */ cancellationSignal, + /* callback = */ dnsResolverCallback, + ) + } catch (e: SecurityException) { + // The app lacks INTERNET permission, or network policy forbids the query. The callback + // won't run, so report the failure here or the state machine never terminates. + taskRunner.newQueue().execute("${question.name} dns", cancelable = false) { callback.onFailure(IOException(e)) } } } - /** One outstanding transport-layer query. */ - private class Query( - question: Question, - ) { - val hostname: String = question.name - val type: Int = question.type - - /** Only the `HTTPS` query reaches [DnsResolver], which is the API that takes a signal. */ - val cancellationSignal = if (type == TYPE_HTTPS) CancellationSignal() else null + /** + * Cancels a query. Only the `HTTPS` query reaches [DnsResolver]; the address lookup runs to + * completion and its result is discarded by [StateMachineDnsCall], which ignores callbacks + * once canceled. + */ + override fun cancel() { + cancellationSignal?.cancel() } } - -/** - * Wraps system-resolved [addresses] in a successful [DnsMessage] so they can flow through - * [StateMachineDnsCall], which only accepts wire-format responses. - */ -private fun Array.toDnsMessage(hostname: String): DnsMessage = - DnsMessage( - id = 0, - // QR = 1 (Response), RCODE = 0 (Success). The state machine only reads the response code. - flags = 0b1___0000__0__0__1__0_000__0000, - questions = listOf(Question(hostname, TYPE_A)), - answers = map { ResourceRecord.IpAddress(name = hostname, timeToLive = 0, address = it) }, - ) +} diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt index a6e16739fc11..3c8e7af5c091 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessage.kt @@ -46,6 +46,25 @@ data class DnsMessage( questions = listOf(question), ) } + + fun response( + id: Short = 0, + questions: List, + answers: List, + ): DnsMessage { + // QR = 1 (Response) + // OPCODE = 0 (standard query) + // RCODE = 0 (success) + // QR OPCODE AA TC RD RA Z RCODE + val flags = 0b1___0000__0__0__0__0_000__0000 + + return DnsMessage( + id = id, + flags = flags, + questions = questions, + answers = answers, + ) + } } } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-StateMachineDnsCall.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-StateMachineDnsCall.kt index 8a8990de1bb6..f068f7fa5348 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-StateMachineDnsCall.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-StateMachineDnsCall.kt @@ -26,18 +26,17 @@ import okhttp3.Protocol import okhttp3.internal.OkHttpInternalApi /** - * An application-layer DNS call that performs multiple transport-layer DNS queries in parallel. - * This delegates to an arbitrary transport like UDP or DNS over HTTPS. + * An application-layer [Dns.Call] that performs multiple transport-layer [DnsQuery]s in parallel. + * This delegates to a query factory for the transport, like UDP or DNS over HTTPS. * * Concurrency * ----------- * * A few things conspire to make concurrency tricky: * - * * 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. + * * Each transport-layer [DnsQuery.Callback]s are executed in parallel. + * * Application layer [Dns.Callback]s must be serialized. + * * We don't want to use locks to guard access to [Dns.Callback] functions. * * Each time we receive data for the callback (in the form of records or an exception), we either * immediately call the callback with that data (on a dispatcher thread), or queue it for the thread @@ -53,14 +52,14 @@ import okhttp3.internal.OkHttpInternalApi * that call is executing. */ @OkHttpInternalApi -class StateMachineDnsCall( +class StateMachineDnsCall( override val request: Dns.Request, - private val transport: Transport, + private val queryFactory: DnsQuery.Factory, private val canceledException: IOException?, private val includeIPv6: Boolean, private val includeServiceMetadata: Boolean, ) : Dns.Call { - private val state = AtomicReference>(State.Idle()) + private val state = AtomicReference(State.Idle()) override fun isCanceled() = state.get().canceled @@ -78,7 +77,7 @@ class StateMachineDnsCall( val queries = questions.map { question -> - transport.newQuery(question) + queryFactory.newQuery(question) } while (true) { @@ -87,7 +86,7 @@ class StateMachineDnsCall( ?: error("already enqueued") val next = - State.Running( + State.Running( canceled = previous.canceled, callback = callback, runningQueries = queries, @@ -97,13 +96,12 @@ class StateMachineDnsCall( for (query in queries) { if (previous.canceled || canceledException != null) { - transport.cancel(query) + query.cancel() } - transport.enqueue( - query = query, + query.enqueue( callback = - object : Transport.Callback { + object : DnsQuery.Callback { override fun onResponse(dnsResponse: DnsMessage) { updateStateAndCallCallbacks( completedQuery = query, @@ -133,7 +131,7 @@ class StateMachineDnsCall( if (previous is State.Running) { for (query in previous.runningQueries) { - transport.cancel(query) + query.cancel() } } return @@ -141,7 +139,7 @@ class StateMachineDnsCall( } private fun updateStateAndCallCallbacks( - completedQuery: Q, + completedQuery: DnsQuery, dnsResponse: DnsMessage, ) { val resourceRecords = @@ -194,7 +192,7 @@ class StateMachineDnsCall( } private tailrec fun updateStateAndCallCallbacks( - completedQuery: Q? = null, + completedQuery: DnsQuery? = null, newRecords: List = listOf(), newException: IOException? = null, lockHeldByThisThread: Boolean = false, @@ -232,7 +230,7 @@ class StateMachineDnsCall( // In such cases, hand off any new work to that other thread and be done. if ((!last && allRecords.isEmpty()) || lockHeldByAnotherThread) { val next = - State.Running( + State.Running( canceled = previous.canceled, callback = previous.callback, runningQueries = newRunningQueries, @@ -288,23 +286,23 @@ class StateMachineDnsCall( } } - private sealed interface State { + private sealed interface State { val canceled: Boolean class Idle( override val canceled: Boolean = false, - ) : State { + ) : State { override fun cancel() = Idle(canceled = true) } - class Running( + class Running( override val canceled: Boolean, val callback: Dns.Callback, val lockHeld: Boolean = false, - val runningQueries: List, + val runningQueries: List, val pendingRecords: List = listOf(), val pendingExceptions: List = listOf(), - ) : State { + ) : State { init { check(pendingRecords.isEmpty() || lockHeld) } @@ -322,28 +320,11 @@ class StateMachineDnsCall( class Complete( override val canceled: Boolean, - ) : State { + ) : State { override fun cancel() = Idle(canceled = true) } - fun cancel(): State - } - - interface Transport { - fun newQuery(question: Question): Q - - fun enqueue( - query: Q, - callback: Callback, - ) - - fun cancel(query: Q) - - interface Callback { - fun onFailure(e: IOException) - - fun onResponse(dnsResponse: DnsMessage) - } + fun cancel(): State } } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-CachingTransport.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/DnsCache.kt similarity index 57% rename from okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-CachingTransport.kt rename to okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/DnsCache.kt index 647e68e8e690..15c39192e090 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-CachingTransport.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/DnsCache.kt @@ -26,11 +26,11 @@ import kotlin.time.ExperimentalTime import kotlin.time.TimeSource import okhttp3.internal.OkHttpInternalApi import okhttp3.internal.concurrent.TaskRunner -import okhttp3.internal.dns.StateMachineDnsCall.Transport /** - * A DNS transport that caches responses according to their [ResourceRecord.timeToLive], bounded by - * a user-supplied minimum and maximum cache duration. + * A DNS query cache that stores responses according to their [ResourceRecord.timeToLive]. Each + * entry's lifetime is bounded between [minimumTimeToLive] and [maximumTimeToLive]. The cache's size + * is bounded by a [maxEntryCount]. * * Cache Hits * ---------- @@ -66,16 +66,15 @@ import okhttp3.internal.dns.StateMachineDnsCall.Transport */ @OkHttpInternalApi @OptIn(ExperimentalTime::class) // We know Clock and Instant will be stable in Kotlin 2.3. -class CachingTransport( +class DnsCache( private val taskRunner: TaskRunner, - private val delegate: Transport, - private val timeSource: TimeSource.WithComparableMarks, + 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, maxEntryCount: Int = 1000, -) : Transport> { +) { private val cache = object : MemoryCache( timeSource = timeSource, @@ -83,7 +82,7 @@ class CachingTransport( ) { override fun lastRequestedAt( now: Time, - value: CachingTransport.Entry, + value: Entry, ): Time? { val state = value.state.get() @@ -104,146 +103,110 @@ class CachingTransport( require(revalidateBeforeExpire >= 0.seconds) } - override fun newQuery(question: Question): Query { - val entry = cache.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 = 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 + fun wrap(delegate: DnsQuery.Factory) = + DnsQuery.Factory { question -> + val entry = cache.computeIfAbsent(question) { Entry() } + CacheQuery(question, delegate, entry) } - } /** - * 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! + * An application-layer DNS query that is served by cached data in [entry] or by a new call to the + * underlying transport via [delegate]. If a new call is made, its result is stored in [entry]. */ - 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")) - } + private inner class CacheQuery( + val question: Question, + val delegate: DnsQuery.Factory, + val entry: Entry, + ) : DnsQuery, + DnsQuery.Callback { + var callback: DnsQuery.Callback? = null - return - } - } + override fun enqueue(callback: DnsQuery.Callback) { + check(this.callback == null) { "already enqueued" } + this.callback = callback - /** A query on this transport. */ - class Query( - val entry: CachingTransport.Entry, - ) { - var callback: Transport.Callback>? = null - } + val now = cache.timeSource.markNow() + while (true) { + val previous = entry.state.get() + val result = previous.result + val inFlightCall = previous.inFlightCall - /** - * 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()) + // We use a cached value unless it's expired. + val useCached = result != null && now < result.expireAt - override fun onFailure(e: IOException) { + // 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 + this) + } + + result == null || now >= result.revalidateAt -> { + InFlightCall( + query = delegate.newQuery(question), + sentAt = now, + queries = if (useCached) listOf() else listOf(this), + ) + } + + else -> { + null + } + }, + ) + + if (!entry.state.compareAndSet(previous, next)) continue // Lost a race, retry. + + if (inFlightCall == null && next.inFlightCall != null) { + next.inFlightCall.query.enqueue(this) + } + + if (useCached) { + taskRunner.newQueue().execute("${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() { while (true) { - val previous = state.get() - val sentAt = previous.inFlightCall!!.sentAt - val revalidateDelay = (failureTimeToLive - revalidateBeforeExpire).coerceAtLeast(0.seconds) + 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 - this + if (newQueries.size == inFlightCall.queries.size) return val next = previous.copy( - inFlightCall = null, - result = - Result.Failure( - exception = e, - revalidateAt = sentAt + revalidateDelay, - expireAt = sentAt + failureTimeToLive, + inFlightCall = + inFlightCall.copy( + queries = newQueries, ), ) - if (!state.compareAndSet(previous, next)) continue // Lost a race, retry. + if (!entry.state.compareAndSet(previous, next)) continue // Lost a race, retry. - val queries = previous.inFlightCall.queries - for (query in queries) { - query.callback!!.onFailure(e) + taskRunner.newQueue().execute("${question.name} dns") { + callback!!.onFailure(IOException("canceled")) } return @@ -252,7 +215,7 @@ class CachingTransport( override fun onResponse(dnsResponse: DnsMessage) { while (true) { - val previous = state.get() + val previous = entry.state.get() val sentAt = previous.inFlightCall!!.sentAt val timeToLive = (dnsResponse.answers.minOfOrNull { it.timeToLive } ?: 0) @@ -271,7 +234,7 @@ class CachingTransport( ), ) - if (!state.compareAndSet(previous, next)) continue // Lost a race, retry. + if (!entry.state.compareAndSet(previous, next)) continue // Lost a race, retry. val queries = previous.inFlightCall.queries for (query in queries) { @@ -281,25 +244,57 @@ class CachingTransport( return } } + + override fun onFailure(e: IOException) { + while (true) { + val previous = entry.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 (!entry.state.compareAndSet(previous, next)) continue // Lost a race, retry. + + val queries = previous.inFlightCall.queries + for (query in queries) { + query.callback!!.onFailure(e) + } + + return + } + } + } + + private class Entry { + val state = AtomicReference(State()) } /** A snapshot of the state of a single entry. */ - data class State( + private data class State( val lastRequestedAt: Time? = null, - val inFlightCall: InFlightCall? = null, + val inFlightCall: InFlightCall? = null, val result: Result? = null, ) /** A call to the underlying transport. */ - data class InFlightCall( - val query: Q, + private data class InFlightCall( + val query: DnsQuery, val sentAt: Time, /** The possibly-empty set of queries to notify when this call is complete. */ - val queries: List>, + val queries: List, ) /** A cached result. */ - sealed interface Result { + private sealed interface Result { val revalidateAt: Time val expireAt: Time diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/DnsQuery.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/DnsQuery.kt new file mode 100644 index 000000000000..a42385b176b4 --- /dev/null +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/DnsQuery.kt @@ -0,0 +1,43 @@ +/* + * 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. + */ +@file:Suppress("ktlint:standard:filename") + +package okhttp3.internal.dns + +import java.io.IOException +import okhttp3.internal.OkHttpInternalApi + +/** + * A transport-layer DNS query for a single record type. + * + * This is different from `Dns.Call` which returns multiple record types. + */ +@OkHttpInternalApi +interface DnsQuery { + fun enqueue(callback: Callback) + + fun cancel() + + interface Callback { + fun onFailure(e: IOException) + + fun onResponse(dnsResponse: DnsMessage) + } + + fun interface Factory { + fun newQuery(question: Question): DnsQuery + } +} diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-LookupDnsCall.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/LookupDnsCall.kt similarity index 100% rename from okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-LookupDnsCall.kt rename to okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/LookupDnsCall.kt diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/MemoryCache.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/MemoryCache.kt index 972fc0432957..6447a43a8849 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/MemoryCache.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/MemoryCache.kt @@ -29,7 +29,7 @@ import kotlin.time.TimeSource * This evicts in big batches because each eviction must traverse the entire cache. */ abstract class MemoryCache( - private val timeSource: TimeSource.WithComparableMarks, + val timeSource: TimeSource.WithComparableMarks, val maxSize: Int, ) { private val entries = ConcurrentHashMap() @@ -42,7 +42,7 @@ abstract class MemoryCache( * Returns the time this value was most recently used, in order to make an eviction decision. This * should return null if this element should be evicted immediately. */ - abstract fun lastRequestedAt( + protected abstract fun lastRequestedAt( now: Time, value: V, ): Time? diff --git a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTest.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTest.kt index d136dc5ec546..2d4feea1f822 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTest.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTest.kt @@ -45,9 +45,9 @@ class StateMachineDnsCallTest { ) 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) + val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = queryFactory.takeQuery("lysine.dev", TYPE_A) query1.respondIpAddresses( addresses = blueIpv6s, @@ -84,7 +84,7 @@ class StateMachineDnsCallTest { caching = true, ) lysineCall0.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, addresses = blueIpv4s, @@ -100,7 +100,7 @@ class StateMachineDnsCallTest { caching = true, ) commonhausCall0.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "commonhaus.org", type = TYPE_A, addresses = greenIpv4s, @@ -143,8 +143,8 @@ class StateMachineDnsCallTest { ) call0.enqueue() - val call0QueryIpv6 = transport.takeQuery("lysine.dev", TYPE_AAAA) - val call0QueryIpv4 = transport.takeQuery("lysine.dev", TYPE_A) + val call0QueryIpv6 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val call0QueryIpv4 = queryFactory.takeQuery("lysine.dev", TYPE_A) call0QueryIpv6.respondIpAddresses( addresses = blueIpv6s, @@ -177,7 +177,7 @@ class StateMachineDnsCallTest { caching = true, ) call0.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 30.seconds, @@ -208,7 +208,7 @@ class StateMachineDnsCallTest { caching = true, ) call2.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, addresses = greenIpv4s, @@ -233,7 +233,7 @@ class StateMachineDnsCallTest { ) call0.enqueue() sleep(30.seconds) - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 30.seconds, @@ -251,7 +251,7 @@ class StateMachineDnsCallTest { caching = true, ) call1.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 30.seconds, @@ -272,7 +272,7 @@ class StateMachineDnsCallTest { caching = true, ) call0.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 1.seconds, @@ -306,7 +306,7 @@ class StateMachineDnsCallTest { caching = true, ) call0.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 100.seconds, @@ -325,7 +325,7 @@ class StateMachineDnsCallTest { caching = true, ) call1.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 1.seconds, @@ -347,8 +347,8 @@ class StateMachineDnsCallTest { ) call0.enqueue() - val call0QueryIpv6 = transport.takeQuery("lysine.dev", TYPE_AAAA) - val call0QueryIpv4 = transport.takeQuery("lysine.dev", TYPE_A) + val call0QueryIpv6 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val call0QueryIpv4 = queryFactory.takeQuery("lysine.dev", TYPE_A) val call1 = newCall( @@ -393,13 +393,13 @@ class StateMachineDnsCallTest { ) call0.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_AAAA, timeToLive = 10.seconds, addresses = blueIpv6s, ) - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 10.seconds, @@ -420,12 +420,12 @@ class StateMachineDnsCallTest { call1.enqueue() assertThat(call1.takeAllRecords().addresses()) .isEqualTo(blueIpv6s + blueIpv4s) - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_AAAA, addresses = greenIpv6s, ) - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, addresses = greenIpv4s, @@ -455,13 +455,13 @@ class StateMachineDnsCallTest { ) call0.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_AAAA, timeToLive = 10.seconds, addresses = blueIpv6s, ) - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, timeToLive = 10.seconds, @@ -484,12 +484,12 @@ class StateMachineDnsCallTest { .isEqualTo(blueIpv6s + blueIpv4s) // Note this doesn't respond to TYPE_AAAA yet. val revalidateQuery0 = - transport.takeQuery( + queryFactory.takeQuery( hostname = "lysine.dev", type = TYPE_AAAA, ) val revalidateQuery1 = - transport.takeQuery( + queryFactory.takeQuery( hostname = "lysine.dev", type = TYPE_A, ) @@ -527,9 +527,9 @@ class StateMachineDnsCallTest { ) 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) + val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = queryFactory.takeQuery("lysine.dev", TYPE_A) query1.respondFailure("boom!") @@ -561,8 +561,8 @@ class StateMachineDnsCallTest { ) call0.enqueue() - val queryIpv6 = transport.takeQuery("lysine.dev", TYPE_AAAA) - val queryIpv4 = transport.takeQuery("lysine.dev", TYPE_A) + val queryIpv6 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val queryIpv4 = queryFactory.takeQuery("lysine.dev", TYPE_A) queryIpv6.respondFailure("boom!") queryIpv4.respondIpAddresses( @@ -599,7 +599,7 @@ class StateMachineDnsCallTest { includeServiceMetadata = false, ) call0.enqueue() - transport + queryFactory .takeQuery("lysine.dev", TYPE_A) .respondFailure("boom!") call0.takeOnFailure("boom!") @@ -614,7 +614,7 @@ class StateMachineDnsCallTest { includeServiceMetadata = false, ) call1.enqueue() - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, addresses = greenIpv4s, @@ -634,7 +634,7 @@ class StateMachineDnsCallTest { includeServiceMetadata = false, ) call0.enqueue() - transport + queryFactory .takeQuery("lysine.dev", TYPE_A) .respondFailure("boom!") call0.takeOnFailure("boom!") @@ -650,7 +650,7 @@ class StateMachineDnsCallTest { ) call1.enqueue() call1.takeOnFailure("boom!") - transport.respondToQuery( + queryFactory.respondToQuery( hostname = "lysine.dev", type = TYPE_A, addresses = blueIpv4s, @@ -686,9 +686,9 @@ class StateMachineDnsCallTest { ) 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) + val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = queryFactory.takeQuery("lysine.dev", TYPE_A) onNextEvent = { query2.respondIpAddresses( @@ -729,12 +729,12 @@ class StateMachineDnsCallTest { 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) + queryFactory.takeCancel("lysine.dev", TYPE_HTTPS) + val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS) + queryFactory.takeCancel("lysine.dev", TYPE_AAAA) + val query1 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + queryFactory.takeCancel("lysine.dev", TYPE_A) + val query2 = queryFactory.takeQuery("lysine.dev", TYPE_A) query0.respondFailure("canceled") query1.respondFailure("canceled") @@ -754,9 +754,9 @@ class StateMachineDnsCallTest { 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) + val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = queryFactory.takeQuery("lysine.dev", TYPE_A) query0.respondFailure("canceled") query1.respondFailure("canceled") @@ -776,9 +776,9 @@ class StateMachineDnsCallTest { ) 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) + val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = queryFactory.takeQuery("lysine.dev", TYPE_A) query1.respondIpAddresses( addresses = blueIpv6s, @@ -789,8 +789,8 @@ class StateMachineDnsCallTest { call.cancel() - transport.takeCancel("lysine.dev", TYPE_HTTPS) - transport.takeCancel("lysine.dev", TYPE_A) + queryFactory.takeCancel("lysine.dev", TYPE_HTTPS) + queryFactory.takeCancel("lysine.dev", TYPE_A) query2.respondIpAddresses( addresses = blueIpv4s, @@ -819,9 +819,9 @@ class StateMachineDnsCallTest { ) 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) + val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS) + val query1 = queryFactory.takeQuery("lysine.dev", TYPE_AAAA) + val query2 = queryFactory.takeQuery("lysine.dev", TYPE_A) query1.respondIpAddresses( addresses = blueIpv6s, diff --git a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTester.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTester.kt index a520b6173730..eaaeca276951 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTester.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/StateMachineDnsCallTester.kt @@ -14,11 +14,8 @@ import kotlin.time.Duration.Companion.seconds 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.DnsMessage.Companion.query -import okhttp3.internal.dns.StateMachineDnsCall.Transport import okhttp3.internal.dns.StateMachineDnsCallTester.CallEvent.OnFailure import okhttp3.internal.dns.StateMachineDnsCallTester.CallEvent.OnRecords import okhttp3.internal.dns.StateMachineDnsCallTester.TransportEvent.QueryCanceled @@ -36,7 +33,7 @@ import okio.ByteString fun testStateMachineDnsCall(block: StateMachineDnsCallTester.() -> Unit) { val tester = StateMachineDnsCallTester() tester.block() - assertThat(tester.transport.events.poll(), "unexpected transport event").isNull() + assertThat(tester.queryFactory.events.poll(), "unexpected transport event").isNull() } class StateMachineDnsCallTester internal constructor() { @@ -45,14 +42,11 @@ class StateMachineDnsCallTester internal constructor() { /** Defend against re-entrant calls. */ private var acceptCallbacks: Boolean = true - val transport = Transport() - private val taskFaker = TaskFaker() - private val cachingTransport = - CachingTransport( + private val dnsCache = + DnsCache( taskRunner = taskFaker.taskRunner, - delegate = transport, timeSource = taskFaker.timeSource, minimumTimeToLive = 10.seconds, maximumTimeToLive = 60.seconds, @@ -61,6 +55,8 @@ class StateMachineDnsCallTester internal constructor() { maxEntryCount = 4, ) + val queryFactory = QueryFactory() + fun newCall( request: Dns.Request, includeIPv6: Boolean = true, @@ -77,11 +73,20 @@ class StateMachineDnsCallTester internal constructor() { taskFaker.advanceUntil(taskFaker.nanoTime + duration.inWholeNanoseconds) } - /** Scriptable transport for testing. */ - inner class Transport : StateMachineDnsCall.Transport { + /** Scriptable query factory for testing. */ + inner class QueryFactory : DnsQuery.Factory { val events = LinkedBlockingDeque() - override fun newQuery(question: Question) = Query(question) + override fun newQuery(question: Question): DnsQuery = + object : DnsQuery { + override fun enqueue(callback: DnsQuery.Callback) { + postEvent(QueryEnqueued(question, callback)) + } + + override fun cancel() { + postEvent(QueryCanceled(question)) + } + } private fun postEvent(e: TransportEvent) { events.put(e) @@ -100,7 +105,7 @@ class StateMachineDnsCallTester internal constructor() { hostname: String, type: Int, ): QueryEnqueued { - val event = transport.takeEvent() as QueryEnqueued + val event = queryFactory.takeEvent() as QueryEnqueued assertThat(event.hostname).isEqualTo(hostname) assertThat(event.type).isEqualTo(type) return event @@ -123,22 +128,11 @@ class StateMachineDnsCallTester internal constructor() { hostname: String, type: Int, ): QueryCanceled { - val event = transport.takeEvent() as QueryCanceled + val event = queryFactory.takeEvent() as QueryCanceled assertThat(event.hostname).isEqualTo(hostname) assertThat(event.type).isEqualTo(type) return event } - - override fun enqueue( - query: Query, - callback: Transport.Callback, - ) { - postEvent(QueryEnqueued(query, callback)) - } - - override fun cancel(query: Query) { - postEvent(QueryCanceled(query)) - } } /** A DNS call for the fake state machine. */ @@ -153,10 +147,10 @@ class StateMachineDnsCallTester internal constructor() { val call = StateMachineDnsCall( request = request, - transport = + queryFactory = when { - caching -> cachingTransport - else -> transport + caching -> dnsCache.wrap(queryFactory) + else -> queryFactory }, canceledException = null, includeIPv6 = includeIPv6, @@ -272,19 +266,15 @@ class StateMachineDnsCallTester internal constructor() { } } - class Query( - val question: Question, - ) - sealed interface TransportEvent { class QueryEnqueued( - val query: Query, - val callback: Transport.Callback, + val question: Question, + val callback: DnsQuery.Callback, ) : TransportEvent { val hostname: String - get() = query.question.name + get() = question.name val type: Int - get() = query.question.type + get() = question.type /** Respond to a [TYPE_HTTPS] query with service metadata. */ fun respondServiceMetadata( @@ -293,16 +283,17 @@ class StateMachineDnsCallTester internal constructor() { echConfigList: ByteString? = null, ) { callback.onResponse( - dnsResponse( - query(query.question), - listOf( - ResourceRecord.Https( - name = query.question.name, - timeToLive = timeToLive.inWholeSeconds.toInt(), - alpnIds = alpnIds, - echConfigList = echConfigList, + DnsMessage.response( + questions = listOf(question), + answers = + listOf( + ResourceRecord.Https( + name = question.name, + timeToLive = timeToLive.inWholeSeconds.toInt(), + alpnIds = alpnIds, + echConfigList = echConfigList, + ), ), - ), ), ) } @@ -318,27 +309,28 @@ class StateMachineDnsCallTester internal constructor() { addresses: List = listOf(), ) { callback.onResponse( - dnsResponse( - query(query.question), - addresses.map { address -> - ResourceRecord.IpAddress( - name = query.question.name, - timeToLive = timeToLive.inWholeSeconds.toInt(), - address = address, - ) - }, + DnsMessage.response( + questions = listOf(question), + answers = + addresses.map { address -> + ResourceRecord.IpAddress( + name = question.name, + timeToLive = timeToLive.inWholeSeconds.toInt(), + address = address, + ) + }, ), ) } } class QueryCanceled( - val query: Query, + val question: Question, ) : TransportEvent { val hostname: String - get() = query.question.name + get() = question.name val type: Int - get() = query.question.type + get() = question.type } }