fix(quic): packet codec hostile-input validation + RFC-correct Initial padding
Round-2 audit found I'd only fixed varint-truncation in Frame.kt, not in
LongHeaderPacket. Plus the datagram padding was wire-illegal AND broke
our own short-header receive path.
LongHeaderPacket.parseAndDecrypt + peekHeader:
- Both tokenLen and length now validated against r.remaining BEFORE
.toInt() truncation. Hostile peer sending length=2^62-1 → return null
instead of OOM/crash.
- dcidLen and scidLen validated 1..20 per RFC 9000 §17.2.
- peekHeader correctly handles Retry packets (no token-length, no
length, no packet-number fields per RFC 9000 §17.2.5) — returns
PeekedHeader with totalLength = bytes.size - offset instead of
reading nonexistent length field as garbage.
Buffer.ensure: doubling loop spins forever if buf.size == 0. coerce
starting newSize to ≥ 8 so the loop always terminates.
QuicConnectionWriter Initial padding (RFC 9000 §14.1):
Previously: built packets, computed datagram size, appended trailing
zero bytes after the last packet to reach 1200. That's wire-illegal
(RFC says PADDING frames inside the encryption envelope) AND breaks
our own short-header receiver because ShortHeaderPacket.parseAndDecrypt
uses bytes.size as packetEnd → AEAD authentication includes the trailing
zeros and fails.
Now: two-pass build. Phase 1 collects ACK+CRYPTO frames per level into
local lists (single takeChunk per level), builds natural-size packets.
If Initial is present and total < 1200, rewinds the Initial PN and
rebuilds with PADDING bytes appended to the encrypted payload. PADDING
is one 0x00 byte per frame so concatenating N zero bytes inside the
payload is N valid PADDING frames — the receiver collapses them to
nothing during decode.
PacketNumberSpaceState.rewindOutboundForRebuild() supports the rebuild.
InMemoryQuicPipe handshake test still passes — verifies the receiver
correctly accepts our new PADDING-frame Initials.
https://claude.ai/code/session_01EC1tfXfap8k8GyKvrxkxZx
This commit is contained in:
@@ -176,7 +176,9 @@ class QuicWriter(
|
||||
|
||||
private fun ensure(more: Int) {
|
||||
if (pos + more > buf.size) {
|
||||
var newSize = buf.size * 2
|
||||
// coerce to a non-zero starting size so the doubling loop terminates
|
||||
// even when initialCapacity was 0.
|
||||
var newSize = (buf.size * 2).coerceAtLeast(8)
|
||||
while (newSize < pos + more) newSize *= 2
|
||||
buf = buf.copyOf(newSize)
|
||||
}
|
||||
|
||||
@@ -55,6 +55,15 @@ class PacketNumberSpaceState {
|
||||
/** Allocate the next outbound packet number. */
|
||||
fun allocateOutbound(): Long = nextPacketNumber++
|
||||
|
||||
/**
|
||||
* Roll back the most recent allocation. Used by the writer when it has
|
||||
* to re-build a packet (e.g. add Initial padding) using the same PN —
|
||||
* the next call to [allocateOutbound] returns the same number again.
|
||||
*/
|
||||
fun rewindOutboundForRebuild() {
|
||||
if (nextPacketNumber > 0) nextPacketNumber--
|
||||
}
|
||||
|
||||
/** Note that an inbound packet was successfully decrypted. */
|
||||
fun observeInbound(
|
||||
packetNumber: Long,
|
||||
|
||||
+65
-34
@@ -61,32 +61,49 @@ fun drainOutbound(
|
||||
return packet
|
||||
}
|
||||
|
||||
val initialPkt = buildInitialPacket(conn, nowMillis)
|
||||
if (initialPkt != null) parts += initialPkt
|
||||
|
||||
val handshakePkt = buildHandshakePacket(conn, nowMillis)
|
||||
if (handshakePkt != null) parts += handshakePkt
|
||||
|
||||
// Drain destructive frame sources into local lists, ONCE.
|
||||
val initialFrames = collectHandshakeLevelFrames(conn, EncryptionLevel.INITIAL, nowMillis)
|
||||
val handshakeFrames = collectHandshakeLevelFrames(conn, EncryptionLevel.HANDSHAKE, nowMillis)
|
||||
val applicationPkt = buildApplicationPacket(conn, nowMillis)
|
||||
if (applicationPkt != null) parts += applicationPkt
|
||||
|
||||
if (parts.isEmpty()) return null
|
||||
val initialState = conn.initial
|
||||
val handshakeState = conn.handshake
|
||||
val initialHasContent = initialFrames != null && initialState.sendProtection != null
|
||||
val handshakeHasContent = handshakeFrames != null && handshakeState.sendProtection != null
|
||||
|
||||
// Coalesce
|
||||
// Build natural-size first.
|
||||
val initialNatural = if (initialHasContent) buildLongHeaderFromFrames(conn, EncryptionLevel.INITIAL, initialFrames!!, padBytes = 0) else null
|
||||
val handshakeNatural = if (handshakeHasContent) buildLongHeaderFromFrames(conn, EncryptionLevel.HANDSHAKE, handshakeFrames!!, padBytes = 0) else null
|
||||
|
||||
val firstPass = listOfNotNull(initialNatural, handshakeNatural, applicationPkt)
|
||||
if (firstPass.isEmpty()) return null
|
||||
|
||||
// RFC 9000 §14.1: client datagrams containing Initial MUST pad to ≥ 1200,
|
||||
// via PADDING frames inside the Initial's encryption envelope.
|
||||
if (initialNatural != null) {
|
||||
var natural = 0
|
||||
for (p in firstPass) natural += p.size
|
||||
if (natural < 1200) {
|
||||
val deficit = 1200 - natural
|
||||
// Rewind the Initial PN — we'll reissue with the same PN and the
|
||||
// same captured frames plus padding.
|
||||
initialState.pnSpace.rewindOutboundForRebuild()
|
||||
val paddedInitial = buildLongHeaderFromFrames(conn, EncryptionLevel.INITIAL, initialFrames!!, padBytes = deficit)
|
||||
return concat(listOfNotNull(paddedInitial, handshakeNatural, applicationPkt))
|
||||
}
|
||||
}
|
||||
return concat(firstPass)
|
||||
}
|
||||
|
||||
private fun concat(parts: List<ByteArray>): ByteArray {
|
||||
var totalLen = 0
|
||||
for (p in parts) totalLen += p.size
|
||||
// Pad the datagram to 1200 bytes if it begins with an Initial packet (RFC 9000 §14).
|
||||
if (initialPkt != null && totalLen < 1200) totalLen = 1200
|
||||
val out = ByteArray(totalLen)
|
||||
var pos = 0
|
||||
for (p in parts) {
|
||||
p.copyInto(out, pos)
|
||||
pos += p.size
|
||||
}
|
||||
// Tail bytes are zero (PADDING frames in the encryption envelope of the last packet).
|
||||
// Note: padding inside the last QUIC packet's encryption is the RFC-correct way; we instead
|
||||
// pad the datagram with zero bytes after the last packet. Servers tolerate this for the
|
||||
// datagram-level minimum size requirement.
|
||||
return out
|
||||
}
|
||||
|
||||
@@ -153,37 +170,51 @@ private fun buildLongHeaderPacket(
|
||||
)
|
||||
}
|
||||
|
||||
private fun buildInitialPacket(
|
||||
conn: QuicConnection,
|
||||
nowMillis: Long,
|
||||
): ByteArray? = buildHandshakeLevelPacket(conn, EncryptionLevel.INITIAL, nowMillis)
|
||||
|
||||
private fun buildHandshakePacket(
|
||||
conn: QuicConnection,
|
||||
nowMillis: Long,
|
||||
): ByteArray? = buildHandshakeLevelPacket(conn, EncryptionLevel.HANDSHAKE, nowMillis)
|
||||
|
||||
private fun buildHandshakeLevelPacket(
|
||||
/**
|
||||
* Drain ACK + CRYPTO frames for [level] into a fresh list. Destructive on
|
||||
* the CRYPTO send buffer, but caller-controlled — call exactly once per
|
||||
* outbound datagram. Returns null if there are no frames to send and the
|
||||
* level has no protection installed.
|
||||
*/
|
||||
private fun collectHandshakeLevelFrames(
|
||||
conn: QuicConnection,
|
||||
level: EncryptionLevel,
|
||||
nowMillis: Long,
|
||||
): ByteArray? {
|
||||
): List<Frame>? {
|
||||
val state = conn.levelState(level)
|
||||
val proto = state.sendProtection ?: return null
|
||||
if (state.sendProtection == null) return null
|
||||
val frames = mutableListOf<Frame>()
|
||||
|
||||
// Pending ACK?
|
||||
state.ackTracker.buildAckFrame(nowMillis, conn.config.ackDelayExponent.toInt())?.let { frames += it }
|
||||
|
||||
// Pending CRYPTO?
|
||||
val cryptoChunk = state.cryptoSend.takeChunk(maxBytes = 1100)
|
||||
if (cryptoChunk != null && cryptoChunk.data.isNotEmpty()) {
|
||||
frames += CryptoFrame(cryptoChunk.offset, cryptoChunk.data)
|
||||
}
|
||||
|
||||
if (frames.isEmpty()) return null
|
||||
return frames
|
||||
}
|
||||
|
||||
val payload = encodeFrames(frames)
|
||||
/**
|
||||
* Build a long-header packet from already-collected frames, with optional
|
||||
* trailing PADDING (0x00) bytes inside the encryption envelope. RFC 9000
|
||||
* §14.1 mandates this for client-Initial datagrams; PADDING is a one-byte
|
||||
* frame so concatenating N zero bytes after the encoded frames yields N
|
||||
* valid PADDING frames, all collapsed to nothing on decode.
|
||||
*/
|
||||
private fun buildLongHeaderFromFrames(
|
||||
conn: QuicConnection,
|
||||
level: EncryptionLevel,
|
||||
frames: List<Frame>,
|
||||
padBytes: Int,
|
||||
): ByteArray {
|
||||
val state = conn.levelState(level)
|
||||
val proto = state.sendProtection!!
|
||||
val basePayload = encodeFrames(frames)
|
||||
val payload =
|
||||
if (padBytes > 0) {
|
||||
ByteArray(basePayload.size + padBytes).also { basePayload.copyInto(it, 0) }
|
||||
} else {
|
||||
basePayload
|
||||
}
|
||||
val pn = state.pnSpace.allocateOutbound()
|
||||
val type = if (level == EncryptionLevel.INITIAL) LongHeaderType.INITIAL else LongHeaderType.HANDSHAKE
|
||||
return LongHeaderPacket.build(
|
||||
|
||||
@@ -144,22 +144,27 @@ object LongHeaderPacket {
|
||||
val packetStart = offset
|
||||
val r = QuicReader(bytes, offset)
|
||||
val first = r.readByte()
|
||||
require((first and 0x80) != 0) { "not a long-header packet" }
|
||||
if ((first and 0x80) == 0) return null // not a long header — silently drop
|
||||
val typeBits = (first ushr 4) and 0x03
|
||||
val type = LongHeaderType.fromTypeBits(typeBits)
|
||||
val version = r.readUint32().toInt()
|
||||
val dcidLen = r.readByte()
|
||||
if (dcidLen !in 0..20) return null // RFC 9000 §17.2 caps CID at 20 bytes
|
||||
val dcidBytes = r.readBytes(dcidLen)
|
||||
val scidLen = r.readByte()
|
||||
if (scidLen !in 0..20) return null
|
||||
val scidBytes = r.readBytes(scidLen)
|
||||
val token =
|
||||
if (type == LongHeaderType.INITIAL) {
|
||||
val tokenLen = r.readVarint().toInt()
|
||||
r.readBytes(tokenLen)
|
||||
val tokenLenRaw = r.readVarint()
|
||||
if (tokenLenRaw < 0L || tokenLenRaw > r.remaining.toLong()) return null
|
||||
r.readBytes(tokenLenRaw.toInt())
|
||||
} else {
|
||||
ByteArray(0)
|
||||
}
|
||||
val length = r.readVarint().toInt()
|
||||
val lengthRaw = r.readVarint()
|
||||
if (lengthRaw < 0L || lengthRaw > r.remaining.toLong()) return null
|
||||
val length = lengthRaw.toInt()
|
||||
val pnOffset = r.position
|
||||
if (pnOffset + length > bytes.size) return null
|
||||
// Sample for HP starts at pnOffset + 4.
|
||||
@@ -233,18 +238,32 @@ object LongHeaderPacket {
|
||||
val type = LongHeaderType.fromTypeBits(typeBits)
|
||||
val version = r.readUint32().toInt()
|
||||
val dcidLen = r.readByte()
|
||||
if (dcidLen !in 0..20) return null
|
||||
val dcid = r.readBytes(dcidLen)
|
||||
val scidLen = r.readByte()
|
||||
if (scidLen !in 0..20) return null
|
||||
val scid = r.readBytes(scidLen)
|
||||
val tokenLen =
|
||||
if (type == LongHeaderType.INITIAL) {
|
||||
r.readVarint().toInt()
|
||||
} else {
|
||||
0
|
||||
}
|
||||
if (type == LongHeaderType.INITIAL) r.skip(tokenLen)
|
||||
val length = r.readVarint().toInt()
|
||||
val total = r.position - offset + length
|
||||
// Retry packets have NO token-length, NO length, NO packet-number fields —
|
||||
// they consist of (header)(retry_token)(16-byte integrity tag). The caller
|
||||
// routes them via [RetryPacket.parse] separately. We surface them here only
|
||||
// so the caller knows the total length is the rest of the input buffer.
|
||||
if (type == LongHeaderType.RETRY) {
|
||||
return PeekedHeader(
|
||||
type = type,
|
||||
version = version,
|
||||
dcid = ConnectionId(dcid),
|
||||
scid = ConnectionId(scid),
|
||||
totalLength = bytes.size - offset,
|
||||
)
|
||||
}
|
||||
if (type == LongHeaderType.INITIAL) {
|
||||
val tokenLenRaw = r.readVarint()
|
||||
if (tokenLenRaw < 0L || tokenLenRaw > r.remaining.toLong()) return null
|
||||
r.skip(tokenLenRaw.toInt())
|
||||
}
|
||||
val lengthRaw = r.readVarint()
|
||||
if (lengthRaw < 0L || lengthRaw > r.remaining.toLong()) return null
|
||||
val total = r.position - offset + lengthRaw.toInt()
|
||||
return PeekedHeader(type, version, ConnectionId(dcid), ConnectionId(scid), total)
|
||||
} catch (_: QuicCodecException) {
|
||||
return null
|
||||
|
||||
Reference in New Issue
Block a user