Merge pull request #2714 from vitorpamplona/claude/dns-resolver-okhtp-XCJWi

Add SurgeDns: burst-tolerant DNS resolver for OkHttp
This commit is contained in:
Vitor Pamplona
2026-05-03 17:28:52 -04:00
committed by GitHub
10 changed files with 1039 additions and 5 deletions
@@ -57,6 +57,8 @@ import com.vitorpamplona.amethyst.service.okhttp.DualHttpClientManager
import com.vitorpamplona.amethyst.service.okhttp.DualHttpClientManagerForRelays
import com.vitorpamplona.amethyst.service.okhttp.EncryptionKeyCache
import com.vitorpamplona.amethyst.service.okhttp.OkHttpWebSocket
import com.vitorpamplona.amethyst.service.okhttp.SurgeDns
import com.vitorpamplona.amethyst.service.okhttp.SurgeDnsStore
import com.vitorpamplona.amethyst.service.playback.diskCache.VideoCache
import com.vitorpamplona.amethyst.service.playback.diskCache.VideoCacheFactory
import com.vitorpamplona.amethyst.service.playback.pip.BackgroundMedia
@@ -201,6 +203,16 @@ class AppModules(
// Key cache service to download and decrypt encrypted files before caching them.
val keyCache = EncryptionKeyCache()
// Concurrent, caching DNS resolver shared by every OkHttp client built below — a host
// resolved for an image fetch is reused when a relay handshake or NIP-05 lookup hits the
// same host.
val surgeDns = SurgeDns()
// Persists [surgeDns]'s positive cache across process restarts so cold starts don't pay
// ~700 sync getaddrinfo calls. Restored entries fall through to the stale-while-revalidate
// path on first lookup.
val dnsStore = SurgeDnsStore(appContext, surgeDns)
// manages all the other connections separately from relays.
val okHttpClients =
DualHttpClientManager(
@@ -209,6 +221,7 @@ class AppModules(
isMobileDataProvider = connManager.isMobileOrNull,
keyCache = keyCache,
scope = applicationIOScope,
dns = surgeDns,
)
// Offers easy methods to know when connections are happening through Tor or not
@@ -290,6 +303,7 @@ class AppModules(
proxyPortProvider = torManager.activePortOrNull,
isMobileDataProvider = connManager.isMobileOrNull,
scope = applicationIOScope,
dns = surgeDns,
)
// Connects the INostrClient class with okHttp
@@ -501,6 +515,22 @@ class AppModules(
fun initiate(appContext: Context) {
Thread.setDefaultUncaughtExceptionHandler(UnexpectedCrashSaver(crashReportCache, applicationIOScope))
// Restore the persisted DNS cache before any networking starts. Lookups that fire
// before this completes fall through to the sync resolver path (existing behavior);
// once restored, every previously-seen host hits the stale-while-revalidate path
// instead of blocking on getaddrinfo.
applicationIOScope.launch {
dnsStore.load()
}
// Periodically flush the DNS cache. Saves are skipped when nothing has changed.
applicationIOScope.launch {
while (true) {
delay(5 * 60 * 1000L)
dnsStore.save()
}
}
applicationIOScope.launch {
// loads main account quickly.
LocalPreferences.loadAccountConfigFromEncryptedStorage()
@@ -565,12 +595,17 @@ class AppModules(
BackgroundMedia.removeBackgroundControllerAndReleaseIt()
PlaybackServiceClient.shutdown()
alwaysOnNotificationServiceManager.stop()
// Best-effort flush before the scope is cancelled. Android rarely calls onTerminate in
// production, but when it does we get one last chance to persist the cache.
runCatching { dnsStore.save() }
applicationIOScope.cancel("Application onTerminate $appContext")
accountsCache.clear()
}
fun trim() {
applicationIOScope.launch {
// Backgrounding is a natural moment to flush the DNS cache.
dnsStore.save()
val loggedIn = accountsCache.accounts.value.values
trimmingService.run(loggedIn, LocalPreferences.allSavedAccounts())
}
@@ -0,0 +1,55 @@
/*
* Copyright (c) 2025 Vitor Pamplona
*
* Permission is hereby granted, free of charge, to any person obtaining a copy of
* this software and associated documentation files (the "Software"), to deal in
* the Software without restriction, including without limitation the rights to use,
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
* Software, and to permit persons to whom the Software is furnished to do so,
* subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
*/
package com.vitorpamplona.amethyst.service.okhttp
import okhttp3.Call
import okhttp3.EventListener
import java.io.IOException
/**
* Drops a host's [SurgeDns] entry whenever an OkHttp call to it fails outright. We hook
* `callFailed` (final-stage signal after OkHttp has tried every address from the DNS lookup)
* rather than per-attempt `connectFailed`: a multi-A-record host with one bad IP fires the
* latter while OkHttp recovers via the next address, and we don't want to invalidate the cache
* just because the first IP was unreachable.
*
* Used by the relay client. The media path uses [MediaCallEventListener], which folds the same
* invalidation into its `finish` method alongside its existing timing logging.
*/
class DnsInvalidatingEventListener(
private val dns: SurgeDns,
) : EventListener() {
override fun callFailed(
call: Call,
ioe: IOException,
) {
dns.invalidate(call.request().url.host)
}
/** Per-client factory. The listener is stateless, so the same instance serves every call. */
class Factory(
dns: SurgeDns,
) : EventListener.Factory {
private val listener = DnsInvalidatingEventListener(dns)
override fun create(call: Call): EventListener = listener
}
}
@@ -38,8 +38,9 @@ class DualHttpClientManager(
isMobileDataProvider: StateFlow<Boolean?>,
keyCache: EncryptionKeyCache,
scope: CoroutineScope,
dns: SurgeDns,
) : IHttpClientManager {
val factory = OkHttpClientFactory(keyCache, userAgent)
val factory = OkHttpClientFactory(keyCache, userAgent, dns)
val defaultHttpClient: StateFlow<OkHttpClient> =
combine(proxyPortProvider, isMobileDataProvider) { proxy, mobile ->
@@ -35,8 +35,9 @@ class DualHttpClientManagerForRelays(
proxyPortProvider: StateFlow<Int?>,
isMobileDataProvider: StateFlow<Boolean?>,
scope: CoroutineScope,
dns: SurgeDns,
) : IHttpClientManager {
val factory = OkHttpClientFactoryForRelays(userAgent)
val factory = OkHttpClientFactoryForRelays(userAgent, dns)
val defaultHttpClient: StateFlow<OkHttpClient> =
combine(proxyPortProvider, isMobileDataProvider) { proxy, mobile ->
@@ -44,6 +44,7 @@ import java.net.Proxy
class MediaCallEventListener(
private val dispatcher: Dispatcher,
private val connectionPool: ConnectionPool,
private val dns: SurgeDns,
) : EventListener() {
private var callStartNanos = 0L
private var dnsStartNanos = 0L
@@ -122,6 +123,16 @@ class MediaCallEventListener(
call: Call,
error: IOException?,
) {
val host = call.request().url.host
// The whole call failed (after OkHttp exhausted every address from the DNS lookup).
// Drop the cached entry so the next attempt re-resolves instead of trying the same
// dead IPs for up to 24h. Per-attempt connectFailed isn't enough — a multi-A-record
// host can have one bad IP and OkHttp will recover by trying the next one.
if (error != null) {
dns.invalidate(host)
}
val totalMs = (System.nanoTime() - callStartNanos) / 1_000_000
val isSlow = totalMs >= SLOW_CALL_THRESHOLD_MS
val wasQueued = queuedAtStart > 0
@@ -129,7 +140,6 @@ class MediaCallEventListener(
if (error == null && !isSlow && !wasQueued && !isDebug) return
val ttfbMs = if (responseHeadersNanos > 0) (responseHeadersNanos - callStartNanos) / 1_000_000 else -1L
val host = call.request().url.host
val reuseTag = if (connectionReused) "reused" else "new"
val msg =
@@ -169,6 +179,7 @@ class MediaCallEventListener(
class MediaCallEventListenerFactory(
private val dispatcher: Dispatcher,
private val connectionPool: ConnectionPool,
private val dns: SurgeDns,
) : EventListener.Factory {
override fun create(call: Call): EventListener = MediaCallEventListener(dispatcher, connectionPool)
override fun create(call: Call): EventListener = MediaCallEventListener(dispatcher, connectionPool, dns)
}
@@ -35,6 +35,7 @@ import java.util.concurrent.TimeUnit
class OkHttpClientFactory(
keyCache: EncryptionKeyCache,
val userAgent: String,
private val dns: SurgeDns,
) {
// val logging = LoggingInterceptor()
val keyDecryptor = EncryptedBlobInterceptor(keyCache)
@@ -63,7 +64,8 @@ class OkHttpClientFactory(
.Builder()
.dispatcher(dispatcher)
.connectionPool(connectionPool)
.eventListenerFactory(MediaCallEventListenerFactory(dispatcher, connectionPool))
.dns(dns)
.eventListenerFactory(MediaCallEventListenerFactory(dispatcher, connectionPool, dns))
.followRedirects(true)
.followSslRedirects(true)
.addInterceptor(DefaultContentTypeInterceptor(userAgent))
@@ -29,6 +29,7 @@ import java.time.Duration
class OkHttpClientFactoryForRelays(
userAgent: String,
private val dns: SurgeDns,
) {
companion object {
// by picking a random proxy port, the connection will fail as it should.
@@ -55,6 +56,8 @@ class OkHttpClientFactoryForRelays(
OkHttpClient
.Builder()
.dispatcher(myDispatcher)
.dns(dns)
.eventListenerFactory(DnsInvalidatingEventListener.Factory(dns))
.followRedirects(true)
.followSslRedirects(true)
.addInterceptor(DefaultContentTypeInterceptor(userAgent))
@@ -0,0 +1,294 @@
/*
* Copyright (c) 2025 Vitor Pamplona
*
* Permission is hereby granted, free of charge, to any person obtaining a copy of
* this software and associated documentation files (the "Software"), to deal in
* the Software without restriction, including without limitation the rights to use,
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
* Software, and to permit persons to whom the Software is furnished to do so,
* subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
*/
package com.vitorpamplona.amethyst.service.okhttp
import okhttp3.Dns
import java.net.InetAddress
import java.net.UnknownHostException
import java.util.Locale
import java.util.concurrent.CompletableFuture
import java.util.concurrent.ConcurrentHashMap
import java.util.concurrent.ExecutionException
import java.util.concurrent.Executor
import java.util.concurrent.Executors
import java.util.concurrent.RejectedExecutionException
import java.util.concurrent.ThreadLocalRandom
import java.util.concurrent.TimeUnit
import java.util.concurrent.atomic.AtomicBoolean
/**
* Concurrent, caching, stale-while-revalidate DNS resolver for OkHttp.
*
* Tuned for an Amethyst-shaped workload: ~700 relays plus a small set of media/profile/NIP-05
* hosts that reappear constantly, whose IPs change on the order of days.
*
* Properties:
*
* 1. **Lock-free reads.** Cache is a [ConcurrentHashMap]. The hot path (cache hit) takes no
* locks, so 700 concurrent relay reconnects fan out across OkHttp dispatcher threads
* instead of serializing through one monitor.
* 2. **Single-flight coalescing.** N concurrent threads asking for the same host share one
* upstream `getaddrinfo`. The leader resolves; followers block on the same future.
* 3. **Stale-while-revalidate.** Once a host has been resolved, recurring lookups *never*
* block on `getaddrinfo` again. After the soft TTL expires, we return the previous answer
* immediately and kick a background refresh. The refresh is coalesced through the same
* `inflight` map, so a burst of stale reads triggers one refresh per host. If the refresh
* fails (`UnknownHostException`), the cache entry is demoted to negative, so the next
* caller gets a fresh failure rather than forever-wrong stale IPs.
* 4. **Negative entries do *not* serve stale.** When an `UnknownHostException` cache entry
* expires, the next call goes through the synchronous path. We want transient failures to
* recover quickly, not keep returning stale failures.
* 5. **Generous positive TTL with jitter.** Defaults to 24h plus up to 24h of random jitter,
* so each entry expires somewhere in [base, base + jitter]. This breaks the synchronized
* herd that would otherwise form when ~700 relay reconnects all populate the cache in the
* same second.
* 6. **Persistable across restarts.** [snapshot] and [restore] expose the positive cache as a
* plain data list so a companion store can survive process death. Wall-clock millis are
* used for expiries (not [System.nanoTime], which resets per process) so a restored entry
* keeps its remaining lifetime.
*
* Remaining blocking points: the very first lookup of a host blocks on `getaddrinfo`
* (unavoidable there's nothing to serve stale yet), and followers waiting on that first
* lookup block on `future.get()`. Background refreshes never block any caller.
*/
class SurgeDns(
private val delegate: Dns = Dns.SYSTEM,
private val maxEntries: Int = 2000,
private val positiveTtlMs: Long = TimeUnit.HOURS.toMillis(24),
private val positiveTtlJitterMs: Long = TimeUnit.HOURS.toMillis(24),
private val negativeTtlMs: Long = TimeUnit.SECONDS.toMillis(10),
private val refreshExecutor: Executor = DEFAULT_REFRESH_EXECUTOR,
) : Dns {
private val cache = ConcurrentHashMap<String, Entry>()
private val inflight = ConcurrentHashMap<String, CompletableFuture<List<InetAddress>>>()
private val dirty = AtomicBoolean(false)
override fun lookup(hostname: String): List<InetAddress> {
// DNS hostnames are case-insensitive; canonicalize so "Example.com" and "example.com"
// share a cache entry. OkHttp normally hands us lowercase, but custom callers may not.
val key = hostname.lowercase(Locale.ROOT)
cache[key]?.let { entry ->
if (entry.expiresAtMillis > System.currentTimeMillis()) {
return entry.unwrap(key)
}
// Soft-expired positive entry: serve stale, refresh in background. Negative
// entries fall through to the sync path so transient failures recover quickly.
if (entry.addresses.isNotEmpty()) {
scheduleBackgroundRefresh(key)
return entry.addresses
}
}
val newFuture = CompletableFuture<List<InetAddress>>()
val existing = inflight.putIfAbsent(key, newFuture)
return if (existing == null) resolveAsLeader(key, newFuture) else awaitFollower(key, existing)
}
private fun scheduleBackgroundRefresh(host: String) {
val refreshFuture = CompletableFuture<List<InetAddress>>()
// Coalesce: if a sync lookup or a prior refresh is already in flight, skip.
if (inflight.putIfAbsent(host, refreshFuture) != null) return
try {
// The caller already got the stale answer; refresh failures are recorded on the
// future and (for UnknownHostException) demoted in the cache by lookupAndCache.
refreshExecutor.execute {
runCatching { resolveAsLeader(host, refreshFuture) }
}
} catch (_: RejectedExecutionException) {
inflight.remove(host, refreshFuture)
}
}
private fun resolveAsLeader(
host: String,
future: CompletableFuture<List<InetAddress>>,
): List<InetAddress> {
try {
// Re-check after claiming leadership: a peer may have refreshed the cache between
// our miss and our putIfAbsent. Skips a duplicate getaddrinfo in that race.
val freshEntry = cache[host]?.takeIf { it.expiresAtMillis > System.currentTimeMillis() }
val addresses = freshEntry?.unwrap(host) ?: lookupAndCache(host)
future.complete(addresses)
return addresses
} catch (e: Throwable) {
future.completeExceptionally(e)
throw e
} finally {
inflight.remove(host, future)
}
}
private fun lookupAndCache(host: String): List<InetAddress> =
try {
val addresses = delegate.lookup(host).ifEmpty { throw UnknownHostException(host) }
putPositive(host, addresses)
addresses
} catch (e: UnknownHostException) {
putNegative(host)
throw e
}
private fun awaitFollower(
host: String,
future: CompletableFuture<List<InetAddress>>,
): List<InetAddress> =
try {
future.get().ifEmpty { throw UnknownHostException(host) }
} catch (e: ExecutionException) {
val cause = e.cause
if (cause is UnknownHostException) throw cause
throw UnknownHostException(host).apply { if (cause != null) initCause(cause) }
} catch (e: InterruptedException) {
Thread.currentThread().interrupt()
throw UnknownHostException(host).apply { initCause(e) }
}
private fun putPositive(
host: String,
addresses: List<InetAddress>,
) {
cache[host] = Entry(addresses, positiveExpiry())
dirty.set(true)
purgeExpiredIfOverCap()
}
private fun putNegative(host: String) {
// Negative entries are never persisted, so they don't dirty the cache.
cache[host] = Entry(emptyList(), negativeExpiry())
purgeExpiredIfOverCap()
}
/**
* Jittered expiry so a burst of co-written entries (e.g. ~700 relay reconnects at app start)
* doesn't all expire at the same instant 24h later.
*/
private fun positiveExpiry(): Long {
val jitter = if (positiveTtlJitterMs > 0L) ThreadLocalRandom.current().nextLong(positiveTtlJitterMs) else 0L
return System.currentTimeMillis() + positiveTtlMs + jitter
}
private fun negativeExpiry(): Long = System.currentTimeMillis() + negativeTtlMs
private fun purgeExpiredIfOverCap() {
if (cache.size <= maxEntries) return
val now = System.currentTimeMillis()
val it = cache.entries.iterator()
while (it.hasNext()) {
if (it.next().value.expiresAtMillis <= now) it.remove()
}
}
/** Drop all cached entries. Call when the network changes (e.g. WiFi <-> mobile). */
fun invalidate() {
cache.clear()
dirty.set(true)
}
/** Drop a single host's cached entry. Call when a connection to it fails. */
fun invalidate(hostname: String) {
val key = hostname.lowercase(Locale.ROOT)
if (cache.remove(key) != null) dirty.set(true)
}
/**
* Snapshot of the positive cache for persistence. Negative entries and expired entries are
* omitted. The expiry timestamps are wall-clock millis (epoch), so they remain meaningful
* across process restarts.
*/
fun snapshot(): List<DnsCacheRecord> {
val now = System.currentTimeMillis()
val out = ArrayList<DnsCacheRecord>(cache.size)
for ((host, entry) in cache) {
if (entry.addresses.isEmpty()) continue
if (entry.expiresAtMillis <= now) continue
// hostAddress is platform-typed (String!); mapNotNull narrows to String and stays
// defensive against the rare case where it might be null.
val ips = entry.addresses.mapNotNull { it.hostAddress }
if (ips.isNotEmpty()) {
out += DnsCacheRecord(host, ips, entry.expiresAtMillis)
}
}
return out
}
/**
* Restore from a previously taken snapshot. Entries already expired on disk are dropped.
* Existing in-memory entries are NOT overwritten (so a fresh lookup that completes before
* restore lands keeps its newer answer).
*/
fun restore(records: List<DnsCacheRecord>) {
val now = System.currentTimeMillis()
for (record in records) {
if (record.expiresAtMillis <= now) continue
val addresses =
record.addresses.mapNotNull { literal ->
// getByName on a numeric literal parses without doing DNS.
runCatching { InetAddress.getByName(literal) }.getOrNull()
}
if (addresses.isNotEmpty()) {
val key = record.hostname.lowercase(Locale.ROOT)
cache.putIfAbsent(key, Entry(addresses, record.expiresAtMillis))
}
}
}
/** True if the cache has changed since the last [tryClearDirty] / [markDirty] / construction. */
fun isDirty(): Boolean = dirty.get()
/**
* Atomically clears the dirty flag if it was set. Returns `true` if the caller observed the
* flag in the dirty state (and is now responsible for persisting). Use this to bracket a save
* clearing **before** taking a snapshot ensures any concurrent [putPositive] re-marks dirty
* and is captured by the next save instead of being silently lost.
*/
fun tryClearDirty(): Boolean = dirty.compareAndSet(true, false)
/** Re-marks the cache dirty. For a store to call when its persist attempt fails. */
fun markDirty() {
dirty.set(true)
}
private class Entry(
val addresses: List<InetAddress>,
val expiresAtMillis: Long,
) {
fun unwrap(host: String): List<InetAddress> = addresses.ifEmpty { throw UnknownHostException(host) }
}
companion object {
// Small fixed pool of daemon threads. Refreshes block on getaddrinfo (~tens of ms),
// so a handful of threads is plenty even when many hosts go stale at once — extra
// refreshes queue up without blocking any caller, since callers always get the
// stale answer instantly.
private val DEFAULT_REFRESH_EXECUTOR: Executor =
Executors.newFixedThreadPool(8) { r ->
Thread(r, "amethyst-dns-refresh").apply { isDaemon = true }
}
}
}
/** Persistable record. Public so [SurgeDnsStore] can serialize it via Jackson. */
data class DnsCacheRecord(
val hostname: String,
val addresses: List<String>,
val expiresAtMillis: Long,
)
@@ -0,0 +1,101 @@
/*
* Copyright (c) 2025 Vitor Pamplona
*
* Permission is hereby granted, free of charge, to any person obtaining a copy of
* this software and associated documentation files (the "Software"), to deal in
* the Software without restriction, including without limitation the rights to use,
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
* Software, and to permit persons to whom the Software is furnished to do so,
* subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
*/
package com.vitorpamplona.amethyst.service.okhttp
import android.content.Context
import com.fasterxml.jackson.module.kotlin.jacksonObjectMapper
import com.fasterxml.jackson.module.kotlin.readValue
import com.vitorpamplona.quartz.utils.Log
/**
* Persists [SurgeDns]'s positive cache to a small `SharedPreferences` blob so the resolver
* starts hot after a process restart.
*
* Cold starts are the resolver's worst case every host is a sync `getaddrinfo`. With a
* persisted snapshot, every previously-seen host falls into the stale-while-revalidate path on
* first lookup: the cached IP is served immediately, and a background refresh updates it. That
* turns ~700 blocking system calls at app start into zero.
*
* The blob is plain (not encrypted) hostnames are already exposed in the user's signed relay
* list, Coil's image cache, and the system resolver's own state. ~700 entries × ~80 bytes
* ~55 KB of JSON.
*/
class SurgeDnsStore(
private val context: Context,
private val dns: SurgeDns,
) {
private val prefs by lazy { context.getSharedPreferences(PREFS_NAME, Context.MODE_PRIVATE) }
/**
* Read the persisted snapshot and merge it into the resolver. Existing in-memory entries
* are preserved (see [SurgeDns.restore]). Safe to call once at app start. Blocking I/O
* call from a background thread.
*/
fun load() {
val json = prefs.getString(KEY_SNAPSHOT, null) ?: return
val records =
try {
MAPPER.readValue<List<DnsCacheRecord>>(json)
} catch (t: Throwable) {
Log.w(TAG) { "Dropping corrupt DNS cache blob: ${t.message}" }
prefs.edit().remove(KEY_SNAPSHOT).apply()
return
}
// restore() uses putIfAbsent and never marks dirty, so we deliberately do NOT clear the
// dirty flag here — any concurrent put that happened before load completed must still be
// persisted on the next save.
dns.restore(records)
Log.d(TAG) { "Restored ${records.size} DNS cache entries" }
}
/**
* Write the current snapshot to disk if the cache has changed since the last save. Blocking
* I/O call from a background thread.
*/
fun save() {
// Clear BEFORE snapshot so any put racing with the snapshot/write re-marks dirty and is
// captured by the next save instead of being silently lost. compareAndSet ensures two
// concurrent saves don't both proceed.
if (!dns.tryClearDirty()) return
try {
val records = dns.snapshot()
val json = MAPPER.writeValueAsString(records)
prefs.edit().putString(KEY_SNAPSHOT, json).apply()
Log.d(TAG) { "Persisted ${records.size} DNS cache entries" }
} catch (t: Throwable) {
Log.w(TAG) { "Failed to persist DNS cache: ${t.message}" }
// Restore the dirty signal so the next save retries.
dns.markDirty()
}
}
/** Force-clear the on-disk cache. Useful for diagnostics or when the user wipes data. */
fun clear() {
prefs.edit().remove(KEY_SNAPSHOT).apply()
}
companion object {
private const val TAG = "SurgeDnsStore"
private const val PREFS_NAME = "amethyst_dns_cache"
private const val KEY_SNAPSHOT = "dns_cache_v1"
private val MAPPER = jacksonObjectMapper()
}
}
@@ -0,0 +1,531 @@
/*
* Copyright (c) 2025 Vitor Pamplona
*
* Permission is hereby granted, free of charge, to any person obtaining a copy of
* this software and associated documentation files (the "Software"), to deal in
* the Software without restriction, including without limitation the rights to use,
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
* Software, and to permit persons to whom the Software is furnished to do so,
* subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
*/
package com.vitorpamplona.amethyst.service.okhttp
import okhttp3.Dns
import org.junit.Assert.assertEquals
import org.junit.Assert.assertFalse
import org.junit.Assert.assertSame
import org.junit.Assert.assertThrows
import org.junit.Assert.assertTrue
import org.junit.Test
import java.net.InetAddress
import java.net.UnknownHostException
import java.util.concurrent.CountDownLatch
import java.util.concurrent.Executor
import java.util.concurrent.Executors
import java.util.concurrent.RejectedExecutionException
import java.util.concurrent.TimeUnit
import java.util.concurrent.atomic.AtomicInteger
import java.util.concurrent.atomic.AtomicReference
class SurgeDnsTest {
private fun ip(value: String) = InetAddress.getByName(value)
private class CountingDns(
private val responses: Map<String, List<InetAddress>>,
) : Dns {
val callsByHost = mutableMapOf<String, AtomicInteger>()
override fun lookup(hostname: String): List<InetAddress> {
callsByHost.getOrPut(hostname) { AtomicInteger() }.incrementAndGet()
return responses[hostname] ?: throw UnknownHostException(hostname)
}
fun calls(hostname: String): Int = callsByHost[hostname]?.get() ?: 0
}
private class GatedDns(
private val responses: Map<String, List<InetAddress>>,
) : Dns {
val started = CountDownLatch(1)
val release = CountDownLatch(1)
val calls = AtomicInteger()
override fun lookup(hostname: String): List<InetAddress> {
calls.incrementAndGet()
started.countDown()
release.await()
return responses[hostname] ?: throw UnknownHostException(hostname)
}
}
@Test
fun `cache hit avoids second upstream call`() {
val upstream = CountingDns(mapOf("a.example" to listOf(ip("1.2.3.4"))))
val dns = SurgeDns(delegate = upstream)
val first = dns.lookup("a.example")
val second = dns.lookup("a.example")
assertEquals(listOf(ip("1.2.3.4")), first)
assertSame(first, second)
assertEquals(1, upstream.calls("a.example"))
}
@Test
fun `negative cache short-circuits subsequent lookups`() {
val upstream = CountingDns(emptyMap())
val dns = SurgeDns(delegate = upstream)
assertThrows(UnknownHostException::class.java) { dns.lookup("missing.example") }
assertThrows(UnknownHostException::class.java) { dns.lookup("missing.example") }
assertEquals(1, upstream.calls("missing.example"))
}
@Test
fun `expired positive entry serves stale and refreshes in background`() {
val upstream = CountingDns(mapOf("a.example" to listOf(ip("1.2.3.4"))))
val syncRefresh = Executor { it.run() }
val dns =
SurgeDns(
delegate = upstream,
positiveTtlMs = 1,
positiveTtlJitterMs = 0,
negativeTtlMs = 1,
refreshExecutor = syncRefresh,
)
dns.lookup("a.example")
Thread.sleep(20)
// Returns the stale cached value AND triggers a refresh on the synchronous executor.
assertEquals(listOf(ip("1.2.3.4")), dns.lookup("a.example"))
assertEquals(2, upstream.calls("a.example"))
}
@Test
fun `expired negative entry does not stale-while-revalidate`() {
val upstream = CountingDns(emptyMap())
val dns =
SurgeDns(
delegate = upstream,
positiveTtlJitterMs = 0,
negativeTtlMs = 1,
)
assertThrows(UnknownHostException::class.java) { dns.lookup("missing.example") }
Thread.sleep(20)
assertThrows(UnknownHostException::class.java) { dns.lookup("missing.example") }
// Two synchronous calls — failed lookups must retry quickly, not be served stale.
assertEquals(2, upstream.calls("missing.example"))
}
@Test
fun `stale read returns previous IP while refresh updates the cache`() {
val responses = AtomicReference<List<InetAddress>>(listOf(ip("1.2.3.4")))
val calls = AtomicInteger()
val upstream =
Dns { _ ->
calls.incrementAndGet()
responses.get()
}
val syncRefresh = Executor { it.run() }
val dns =
SurgeDns(
delegate = upstream,
positiveTtlMs = 1,
positiveTtlJitterMs = 0,
refreshExecutor = syncRefresh,
)
assertEquals(listOf(ip("1.2.3.4")), dns.lookup("a.example"))
Thread.sleep(20)
// Upstream now returns a new IP. The next lookup should still serve the OLD IP
// immediately, while the synchronous executor performs the refresh inline.
responses.set(listOf(ip("5.6.7.8")))
val stale = dns.lookup("a.example")
assertEquals(listOf(ip("1.2.3.4")), stale)
assertEquals(2, calls.get())
// Cache now holds the refreshed IP — no further upstream calls.
assertEquals(listOf(ip("5.6.7.8")), dns.lookup("a.example"))
assertEquals(2, calls.get())
}
@Test
fun `stale burst on the same host triggers a single refresh`() {
val gated = GatedDns(mapOf("hot.example" to listOf(ip("9.9.9.9"))))
// Pre-populate via a separate, non-gated upstream then swap in the gated one for
// the refresh — easiest way: bootstrap by writing through a delegate that completes
// immediately, then attach the gate to count parallel refresh calls.
val bootstrapCalls = AtomicInteger()
val dynamic =
object : Dns {
@Volatile var useGated = false
override fun lookup(hostname: String): List<InetAddress> =
if (useGated) {
gated.lookup(hostname)
} else {
bootstrapCalls.incrementAndGet()
listOf(ip("9.9.9.9"))
}
}
val pool = Executors.newFixedThreadPool(8)
val dns =
SurgeDns(
delegate = dynamic,
positiveTtlMs = 1,
positiveTtlJitterMs = 0,
refreshExecutor = pool,
)
try {
dns.lookup("hot.example")
assertEquals(1, bootstrapCalls.get())
Thread.sleep(20)
dynamic.useGated = true
// Fan out 16 stale reads concurrently. They all return the stale answer
// immediately; only one refresh should be queued/executing.
val callerPool = Executors.newFixedThreadPool(16)
try {
val results =
(1..16).map {
callerPool.submit<List<InetAddress>> { dns.lookup("hot.example") }
}
results.forEach { assertEquals(listOf(ip("9.9.9.9")), it.get(2, TimeUnit.SECONDS)) }
assertTrue(
"Refresh should have started",
gated.started.await(2, TimeUnit.SECONDS),
)
gated.release.countDown()
// Allow refresh to complete.
Thread.sleep(100)
assertEquals("Only one refresh upstream call", 1, gated.calls.get())
} finally {
callerPool.shutdownNow()
}
} finally {
pool.shutdownNow()
}
}
@Test
fun `concurrent lookups for the same host coalesce to one upstream call`() {
val gated = GatedDns(mapOf("hot.example" to listOf(ip("9.9.9.9"))))
val dns = SurgeDns(delegate = gated)
val pool = Executors.newFixedThreadPool(8)
try {
val results = (1..8).map { pool.submit<List<InetAddress>> { dns.lookup("hot.example") } }
assertTrue(
"Leader should have started the upstream lookup",
gated.started.await(2, TimeUnit.SECONDS),
)
gated.release.countDown()
results.forEach {
assertEquals(listOf(ip("9.9.9.9")), it.get(2, TimeUnit.SECONDS))
}
assertEquals(1, gated.calls.get())
} finally {
pool.shutdownNow()
}
}
@Test
fun `lookups for different hosts run in parallel`() {
val responses = mapOf("a" to listOf(ip("1.1.1.1")), "b" to listOf(ip("2.2.2.2")))
val parallelism = AtomicInteger()
val peak = AtomicInteger()
val release = CountDownLatch(1)
val instrumented =
Dns { hostname ->
val now = parallelism.incrementAndGet()
peak.updateAndGet { maxOf(it, now) }
try {
release.await(2, TimeUnit.SECONDS)
responses[hostname] ?: throw UnknownHostException(hostname)
} finally {
parallelism.decrementAndGet()
}
}
val dns = SurgeDns(delegate = instrumented)
val pool = Executors.newFixedThreadPool(2)
try {
val futureA = pool.submit<List<InetAddress>> { dns.lookup("a") }
val futureB = pool.submit<List<InetAddress>> { dns.lookup("b") }
// Wait briefly for both threads to enter the resolver, then let them out.
val deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(2)
while (peak.get() < 2 && System.nanoTime() < deadline) {
Thread.sleep(5)
}
release.countDown()
futureA.get(2, TimeUnit.SECONDS)
futureB.get(2, TimeUnit.SECONDS)
assertEquals("Both hosts should resolve concurrently", 2, peak.get())
} finally {
pool.shutdownNow()
}
}
@Test
fun `invalidate clears cache so next lookup hits upstream`() {
val upstream = CountingDns(mapOf("a.example" to listOf(ip("1.2.3.4"))))
val dns = SurgeDns(delegate = upstream)
dns.lookup("a.example")
dns.invalidate()
dns.lookup("a.example")
assertEquals(2, upstream.calls("a.example"))
}
@Test
fun `invalidate by host removes only that entry`() {
val upstream =
CountingDns(
mapOf(
"a.example" to listOf(ip("1.2.3.4")),
"b.example" to listOf(ip("5.6.7.8")),
),
)
val dns = SurgeDns(delegate = upstream)
dns.lookup("a.example")
dns.lookup("b.example")
dns.invalidate("a.example")
dns.lookup("a.example")
dns.lookup("b.example")
assertEquals(2, upstream.calls("a.example"))
assertEquals(1, upstream.calls("b.example"))
}
@Test
fun `snapshot includes only fresh positive entries`() {
val upstream =
CountingDns(
mapOf(
"live.example" to listOf(ip("1.2.3.4")),
"missing.example" to emptyList(),
),
)
val dns =
SurgeDns(
delegate = upstream,
positiveTtlMs = 60_000,
positiveTtlJitterMs = 0,
)
dns.lookup("live.example")
runCatching { dns.lookup("missing.example") }
val snapshot = dns.snapshot()
assertEquals(1, snapshot.size)
assertEquals("live.example", snapshot[0].hostname)
assertEquals(listOf("1.2.3.4"), snapshot[0].addresses)
}
@Test
fun `restore replays cached entries without hitting upstream`() {
val upstream = CountingDns(mapOf("relay.example" to listOf(ip("9.9.9.9"))))
val dns = SurgeDns(delegate = upstream)
val expiresAt = System.currentTimeMillis() + 60_000
dns.restore(listOf(DnsCacheRecord("relay.example", listOf("9.9.9.9"), expiresAt)))
assertEquals(listOf(ip("9.9.9.9")), dns.lookup("relay.example"))
assertEquals("Restored entry should serve without upstream", 0, upstream.calls("relay.example"))
}
@Test
fun `restore drops entries already expired on disk`() {
val upstream = CountingDns(mapOf("relay.example" to listOf(ip("9.9.9.9"))))
val dns = SurgeDns(delegate = upstream)
val expiredAt = System.currentTimeMillis() - 1_000
dns.restore(listOf(DnsCacheRecord("relay.example", listOf("9.9.9.9"), expiredAt)))
dns.lookup("relay.example")
assertEquals(1, upstream.calls("relay.example"))
}
@Test
fun `restore does not overwrite a fresh in-memory entry`() {
val upstream = CountingDns(mapOf("relay.example" to listOf(ip("1.1.1.1"))))
val dns =
SurgeDns(
delegate = upstream,
positiveTtlMs = 60_000,
positiveTtlJitterMs = 0,
)
// Live lookup populates with 1.1.1.1.
dns.lookup("relay.example")
// Then a (stale-on-disk-but-not-yet-expired) restore arrives with a different IP.
dns.restore(
listOf(
DnsCacheRecord(
"relay.example",
listOf("9.9.9.9"),
System.currentTimeMillis() + 60_000,
),
),
)
assertEquals("In-memory entry wins", listOf(ip("1.1.1.1")), dns.lookup("relay.example"))
}
@Test
fun `lookup is case-insensitive`() {
val upstream = CountingDns(mapOf("example.com" to listOf(ip("1.2.3.4"))))
val dns = SurgeDns(delegate = upstream)
dns.lookup("Example.COM")
dns.lookup("example.com")
dns.lookup("ExAmPlE.cOm")
assertEquals("Mixed-case lookups should share one cache entry", 1, upstream.calls("example.com"))
}
@Test
fun `invalidate is case-insensitive`() {
val upstream = CountingDns(mapOf("example.com" to listOf(ip("1.2.3.4"))))
val dns = SurgeDns(delegate = upstream)
dns.lookup("example.com")
dns.invalidate("EXAMPLE.com")
dns.lookup("example.com")
assertEquals(2, upstream.calls("example.com"))
}
@Test
fun `restore lowercases hostnames so subsequent lookups hit`() {
val upstream = CountingDns(mapOf("example.com" to listOf(ip("9.9.9.9"))))
val dns = SurgeDns(delegate = upstream)
dns.restore(
listOf(
DnsCacheRecord("Example.COM", listOf("9.9.9.9"), System.currentTimeMillis() + 60_000),
),
)
assertEquals(listOf(ip("9.9.9.9")), dns.lookup("example.com"))
assertEquals("Lookup should hit restored entry without upstream", 0, upstream.calls("example.com"))
}
@Test
fun `dirty flag tracks positive writes`() {
val upstream = CountingDns(mapOf("a.example" to listOf(ip("1.2.3.4"))))
val dns =
SurgeDns(
delegate = upstream,
positiveTtlMs = 60_000,
positiveTtlJitterMs = 0,
)
assertFalse("Fresh resolver is not dirty", dns.isDirty())
dns.lookup("a.example")
assertTrue("First positive write dirties cache", dns.isDirty())
assertTrue("tryClearDirty reports prior dirty state", dns.tryClearDirty())
dns.lookup("a.example") // cache hit, no write
assertFalse("Cache hit does not dirty", dns.isDirty())
assertFalse("tryClearDirty on already-clean returns false", dns.tryClearDirty())
}
@Test
fun `failed lookup does not mark cache dirty`() {
val upstream = CountingDns(emptyMap())
val dns = SurgeDns(delegate = upstream)
assertFalse(dns.isDirty())
runCatching { dns.lookup("missing.example") }
assertFalse("Negative entry must not dirty the cache (it isn't persisted)", dns.isDirty())
}
@Test
fun `failed refresh demotes stale positive entry to negative`() {
val responses = AtomicReference<List<InetAddress>?>(listOf(ip("1.2.3.4")))
val upstream =
Dns { hostname ->
responses.get() ?: throw UnknownHostException(hostname)
}
val syncRefresh = Executor { it.run() }
val dns =
SurgeDns(
delegate = upstream,
positiveTtlMs = 1,
positiveTtlJitterMs = 0,
negativeTtlMs = 60_000,
refreshExecutor = syncRefresh,
)
// Populate, then let it go stale.
dns.lookup("a.example")
Thread.sleep(20)
// Make upstream fail.
responses.set(null)
// Stale read returns the cached value AND triggers a refresh that fails.
assertEquals(listOf(ip("1.2.3.4")), dns.lookup("a.example"))
// Entry should now be negative — next caller gets a fresh failure rather than
// forever-stale wrong IPs.
assertThrows(UnknownHostException::class.java) { dns.lookup("a.example") }
}
@Test
fun `refresh executor rejection cleans up the inflight slot`() {
val upstream = CountingDns(mapOf("a.example" to listOf(ip("1.2.3.4"))))
val rejecting = Executor { throw RejectedExecutionException("test") }
val dns =
SurgeDns(
delegate = upstream,
positiveTtlMs = 1,
positiveTtlJitterMs = 0,
refreshExecutor = rejecting,
)
dns.lookup("a.example")
Thread.sleep(20)
// Stale read tries to schedule a refresh; the executor rejects it. The caller still
// gets the stale answer.
assertEquals(listOf(ip("1.2.3.4")), dns.lookup("a.example"))
// If the rejected future was leaked into inflight, a subsequent cache-miss lookup for
// the same host would block forever in awaitFollower. Force a cache miss via invalidate
// and run the lookup with a hard timeout to verify the slot was freed.
dns.invalidate("a.example")
val pool = Executors.newSingleThreadExecutor()
try {
val result =
pool
.submit<List<InetAddress>> { dns.lookup("a.example") }
.get(2, TimeUnit.SECONDS)
assertEquals(listOf(ip("1.2.3.4")), result)
} finally {
pool.shutdownNow()
}
}
}