Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -105,26 +105,27 @@ class StateMachineDnsCall(
return
}

val queries =
questions.map { question ->
val questionToQuery =
questions.associateWith { question ->

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

TIL The returned map preserves the entry iteration order of the original array.

queryFactory.newQuery(question)
}

val next =
State.Running(
canceled = false,
callback = callback,
runningQueries = queries,
runningQueries = questionToQuery.values.toList(),
)

if (!state.compareAndSet(previous, next)) continue // Lost a race, retry.

for (query in queries) {
for ((question, query) in questionToQuery) {
query.enqueue(
callback =
object : DnsQuery.Callback {
override fun onResponse(dnsResponse: DnsMessage) {
updateStateAndCallCallbacks(
question = question,
completedQuery = query,
dnsResponse = dnsResponse,
)
Expand Down Expand Up @@ -160,6 +161,7 @@ class StateMachineDnsCall(
}

private fun updateStateAndCallCallbacks(
question: Question,
completedQuery: DnsQuery,
dnsResponse: DnsMessage,
) {
Expand All @@ -177,10 +179,21 @@ class StateMachineDnsCall(
)
}

val dnsRecords =
resourceRecords.map { resourceRecord ->
when (resourceRecord) {
is ResourceRecord.Https -> {
val dnsRecords: List<Dns.Record> =
when (question.type) {
TYPE_HTTPS -> {
resourceRecords.mapNotNull { resourceRecord ->
// Discard resource records that don't fit the query.
if (resourceRecord !is ResourceRecord.Https) return@mapNotNull null

// OkHttp doesn't yet implement AliasMode resource records. If any AliasMode record is
// returned, we must ignore ALL returned resource records.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

TODO log something observable?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

I don't think the recipient of the warning is the person who can do anything about it.

There's probably a DNS record linter tool for DNS admins to learn that their records are malformed.

if (resourceRecord.priority == 0) {
return updateStateAndCallCallbacks(
completedQuery = completedQuery,
)
}

Dns.Record.ServiceMetadata(
hostname = resourceRecord.targetName.takeIf { it != "" } ?: request.hostname,
alpnIds =
Expand All @@ -196,14 +209,23 @@ class StateMachineDnsCall(
echConfigList = resourceRecord.echConfigList,
)
}
}

TYPE_A, TYPE_AAAA -> {
resourceRecords.mapNotNull { resourceRecord ->
// Discard resource records that don't fit the query.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

this seems unlikely and worth warning about, but I assume we think it won't happen, so ignore?

Arguably ew should throw away A answers for a AAAA query, but I'm guessing this is a smart cast.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Yeah even if we did warn, it's unlikely the recipient of the warning would be able to do something with it.

if (resourceRecord !is ResourceRecord.IpAddress) return@mapNotNull null

is ResourceRecord.IpAddress -> {
Dns.Record.IpAddress(
hostname = request.hostname,
address = resourceRecord.address,
)
}
}

else -> {
error("unexpected question type")
}
}

updateStateAndCallCallbacks(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1047,4 +1047,83 @@ class StateMachineDnsCallTest {
assertThat(cache.networkCount).isEqualTo(3)
assertThat(cache.hitCount).isEqualTo(0)
}

@Test
fun `alias mode records are ignored`() =
testStateMachineDnsCall {
val call = newCall(request = Dns.Request(hostname = "lysine.dev"))
call.enqueue()

// Priority 0 means 'AliasMode'. We must ignore all SvcParams in AliasMode.
val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS)
query0.respond(
ResourceRecord.Https(
timeToLive = 300,
name = "lysine.dev",
priority = 0,
alpnIds = listOf("h2"),
),
)
queryFactory.respondToQuery(
hostname = "lysine.dev",
type = TYPE_AAAA,
addresses = blueIpv6s,
)
queryFactory.respondToQuery(
hostname = "lysine.dev",
type = TYPE_A,
addresses = blueIpv4s,
)

call.takeOnRecordsIpAddresses(
last = false,
addresses = blueIpv6s,
)
call.takeOnRecordsIpAddresses(
last = true,
addresses = blueIpv4s,
)
}

@Test
fun `service mode records are ignored if any alias mode record is present`() =
testStateMachineDnsCall {
val call = newCall(request = Dns.Request(hostname = "lysine.dev"))
call.enqueue()

// Priority 0 means 'AliasMode'. We must ignore all SvcParams in AliasMode.
val query0 = queryFactory.takeQuery("lysine.dev", TYPE_HTTPS)
query0.respond(
ResourceRecord.Https(
timeToLive = 300,
name = "lysine.dev",
priority = 0,
),
ResourceRecord.Https(
timeToLive = 300,
name = "lysine.dev",
priority = 1,
alpnIds = listOf("h2"),
),
)
queryFactory.respondToQuery(
hostname = "lysine.dev",
type = TYPE_AAAA,
addresses = blueIpv6s,
)
queryFactory.respondToQuery(
hostname = "lysine.dev",
type = TYPE_A,
addresses = blueIpv4s,
)

call.takeOnRecordsIpAddresses(
last = false,
addresses = blueIpv6s,
)
call.takeOnRecordsIpAddresses(
last = true,
addresses = blueIpv4s,
)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -282,18 +282,12 @@ class StateMachineDnsCallTester internal constructor() {
alpnIds: List<String>? = null,
echConfigList: ByteString? = null,
) {
callback.onResponse(
DnsMessage.response(
questions = listOf(question),
answers =
listOf(
ResourceRecord.Https(
name = question.name,
timeToLive = timeToLive.inWholeSeconds.toInt(),
alpnIds = alpnIds,
echConfigList = echConfigList,
),
),
respond(
ResourceRecord.Https(
name = question.name,
timeToLive = timeToLive.inWholeSeconds.toInt(),
alpnIds = alpnIds,
echConfigList = echConfigList,
),
)
}
Expand All @@ -308,17 +302,22 @@ class StateMachineDnsCallTester internal constructor() {
timeToLive: Duration = 300.seconds,
addresses: List<InetAddress> = listOf(),
) {
val resourceRecords =
addresses.map { address ->
ResourceRecord.IpAddress(
name = question.name,
timeToLive = timeToLive.inWholeSeconds.toInt(),
address = address,
)
}
respond(*resourceRecords.toTypedArray())
}

fun respond(vararg resourceRecords: ResourceRecord) {
callback.onResponse(
DnsMessage.response(
questions = listOf(question),
answers =
addresses.map { address ->
ResourceRecord.IpAddress(
name = question.name,
timeToLive = timeToLive.inWholeSeconds.toInt(),
address = address,
)
},
answers = resourceRecords.toList(),
),
)
}
Expand Down
Loading