From 59c51974a32d91729548bb00cee1c6b5e9867a26 Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Sun, 19 Jul 2026 09:41:50 -0400 Subject: [PATCH] Reject malformed messages that have infinite loops --- .../internal/-DnsMessageReader.kt | 18 +++++++++++++----- .../internal/DnsMessageReaderWriterTest.kt | 19 +++++++++++++++++++ 2 files changed, 32 insertions(+), 5 deletions(-) diff --git a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsMessageReader.kt b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsMessageReader.kt index 7e833e35502f..569fb67679cd 100644 --- a/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsMessageReader.kt +++ b/okhttp-dnsoverhttps/src/main/kotlin/okhttp3/dnsoverhttps/internal/-DnsMessageReader.kt @@ -107,8 +107,10 @@ class DnsMessageReader( return result.readUtf8() } - // TODO: don't infinite loop - private tailrec fun BufferedSource.readName(sink: Buffer) { + private tailrec fun BufferedSource.readName( + sink: Buffer, + maxOffset: Int = Int.MAX_VALUE, + ) { while (true) { val labelTypeAndLength = readByte().toUByte().toInt() val labelType = labelTypeAndLength and 0b11000000 @@ -123,10 +125,16 @@ class DnsMessageReader( // Compressed suffix. 0b11_000000 -> { - val offsetLength = (labelLength shl 8) or readByte().toUByte().toInt() + val offset = (labelLength shl 8) or readByte().toUByte().toInt() + + // Pointers may only refer to prior occurrences. + if (offset >= maxOffset) { + throw ProtocolException("malformed DNS message") + } + val offsetSource = sourceOffsetZero.peek() - offsetSource.skip(offsetLength.toLong()) - return offsetSource.readName(sink) + offsetSource.skip(offset.toLong()) + return offsetSource.readName(sink, maxOffset = offset) } 0b01_000000, 0b10_000000 -> { diff --git a/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/internal/DnsMessageReaderWriterTest.kt b/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/internal/DnsMessageReaderWriterTest.kt index b428943fffa1..f2b500338abc 100644 --- a/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/internal/DnsMessageReaderWriterTest.kt +++ b/okhttp-dnsoverhttps/src/test/java/okhttp3/dnsoverhttps/internal/DnsMessageReaderWriterTest.kt @@ -16,9 +16,12 @@ package okhttp3.dnsoverhttps.internal import assertk.assertThat +import assertk.assertions.hasMessage import assertk.assertions.isEqualTo import java.net.InetAddress +import java.net.ProtocolException import kotlin.test.Test +import kotlin.test.assertFailsWith import okio.Buffer import okio.ByteString.Companion.decodeHex @@ -134,6 +137,22 @@ class DnsMessageReaderWriterTest { ) } + @Test + fun `unbounded name compression`() { + val buffer = Buffer() + buffer.write( + "000081800001000100000000066c7973696e65c00c000100010363646ec00c000100010000000000040a141e28" + .decodeHex(), + ) + + val reader = DnsMessageReader(buffer) + val e = + assertFailsWith { + reader.read() + } + assertThat(e).hasMessage("malformed DNS message") + } + private fun assertRoundTrip(message: DnsMessage) { val buffer = Buffer() DnsMessageWriter(buffer).write(message)