diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessageReader.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessageReader.kt index 95e981f07fb5..8c24c3fc080b 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessageReader.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/dns/-DnsMessageReader.kt @@ -208,8 +208,9 @@ class DnsMessageReader( val valueLength = readUShort().toLong() when (key) { SERVICE_PARAMETER_MANDATORY -> { - for (i in 0 until valueLength) { - val serviceParameterKey = readByte() + if (valueLength % 2 != 0L) throw ProtocolException("malformed HTTPS / mandatory") + for (i in 0 until valueLength step 2) { + val serviceParameterKey = readUShort().toInt() if (serviceParameterKey !in SERVICE_PARAMETER_MANDATORY..SERVICE_PARAMETER_IPV6_HINT) { throw ProtocolException("unsupported HTTPS mandatory parameter $serviceParameterKey") } diff --git a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsMessageReaderWriterTest.kt b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsMessageReaderWriterTest.kt index 8cf33e7528f9..dda2b0a35bfe 100644 --- a/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsMessageReaderWriterTest.kt +++ b/okhttp/src/commonTest/kotlin/okhttp3/internal/dns/DnsMessageReaderWriterTest.kt @@ -157,6 +157,31 @@ class DnsMessageReaderWriterTest { assertThat(e).hasMessage("malformed DNS message") } + @Test + fun `rejects mandatory service parameter with unsupported key`() { + // An HTTPS record whose mandatory value is the single 2-octet SvcParamKey 0x0101 (257), which + // OkHttp does not support. RFC 9460 requires such a record be treated as malformed. + val buffer = Buffer() + buffer.write( + "0000818000000001000000000000410001000000000009000100000000020101".decodeHex(), + ) + assertFailsWith { + DnsMessageReader(buffer).read() + } + } + + @Test + fun `rejects odd length mandatory service parameter`() { + // The mandatory value has an odd length, so it can't be a list of 2-octet SvcParamKeys. + val buffer = Buffer() + buffer.write( + "00008180000000010000000000004100010000000000080001000000000103".decodeHex(), + ) + assertFailsWith { + DnsMessageReader(buffer).read() + } + } + private fun assertRoundTrip(message: DnsMessage) { val buffer = Buffer() DnsMessageWriter(buffer).write(message)