authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-12 04:33:20+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log373d8ef26edca9d16111ae41f960a44ead6ea2c8
tree32de189f254ee9b30ef799d48a6e4be1b043e180
parent1530e73648cd9687bbaea3e50da9b2e86d66df0c

std.compress.zstandard: check FSE bitstreams are fully consumed


3 files changed, 32 insertions(+), 16 deletions(-)

lib/std/compress/zstandard/decode/block.zig+21-15
...@@ -391,15 +391,21 @@ pub const DecodeState = struct {...@@ -391,15 +391,21 @@ pub const DecodeState = struct {
391 try self.literal_stream_reader.init(bytes);391 try self.literal_stream_reader.init(bytes);
392 }392 }
393393
394 fn isLiteralStreamEmpty(self: *DecodeState) bool {
395 switch (self.literal_streams) {
396 .one => return self.literal_stream_reader.isEmpty(),
397 .four => return self.literal_stream_index == 3 and self.literal_stream_reader.isEmpty(),
398 }
399 }
400
394 const LiteralBitsError = error{401 const LiteralBitsError = error{
395 BitStreamHasNoStartBit,402 BitStreamHasNoStartBit,
396 UnexpectedEndOfLiteralStream,403 UnexpectedEndOfLiteralStream,
397 };404 };
398 fn readLiteralsBits(405 fn readLiteralsBits(
399 self: *DecodeState,406 self: *DecodeState,
400 comptime T: type,
401 bit_count_to_read: usize,407 bit_count_to_read: usize,
402 ) LiteralBitsError!T {408 ) LiteralBitsError!u16 {
403 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {409 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {
404 if (self.literal_streams == .four and self.literal_stream_index < 3) {410 if (self.literal_streams == .four and self.literal_stream_index < 3) {
405 try self.nextLiteralMultiStream();411 try self.nextLiteralMultiStream();
...@@ -461,7 +467,7 @@ pub const DecodeState = struct {...@@ -461,7 +467,7 @@ pub const DecodeState = struct {
461 while (i < len) : (i += 1) {467 while (i < len) : (i += 1) {
462 var prefix: u16 = 0;468 var prefix: u16 = 0;
463 while (true) {469 while (true) {
464 const new_bits = self.readLiteralsBits(u16, bit_count_to_read) catch |err| {470 const new_bits = self.readLiteralsBits(bit_count_to_read) catch |err| {
465 return err;471 return err;
466 };472 };
467 prefix <<= bit_count_to_read;473 prefix <<= bit_count_to_read;
...@@ -533,7 +539,7 @@ pub const DecodeState = struct {...@@ -533,7 +539,7 @@ pub const DecodeState = struct {
533 while (i < len) : (i += 1) {539 while (i < len) : (i += 1) {
534 var prefix: u16 = 0;540 var prefix: u16 = 0;
535 while (true) {541 while (true) {
536 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read);542 const new_bits = try self.readLiteralsBits(bit_count_to_read);
537 prefix <<= bit_count_to_read;543 prefix <<= bit_count_to_read;
538 prefix |= new_bits;544 prefix |= new_bits;
539 bits_read += bit_count_to_read;545 bits_read += bit_count_to_read;
...@@ -659,13 +665,10 @@ pub fn decodeBlock(...@@ -659,13 +665,10 @@ pub fn decodeBlock(
659 sequence_size_limit -= decompressed_size;665 sequence_size_limit -= decompressed_size;
660 }666 }
661667
662 if (bit_stream.bit_reader.bit_count != 0) {668 if (!bit_stream.isEmpty()) {
663 return error.MalformedCompressedBlock;669 return error.MalformedCompressedBlock;
664 }670 }
665
666 bytes_read += bit_stream_bytes.len;
667 }671 }
668 if (bytes_read != block_size) return error.MalformedCompressedBlock;
669672
670 if (decode_state.literal_written_count < literals.header.regenerated_size) {673 if (decode_state.literal_written_count < literals.header.regenerated_size) {
671 const len = literals.header.regenerated_size - decode_state.literal_written_count;674 const len = literals.header.regenerated_size - decode_state.literal_written_count;
...@@ -675,7 +678,9 @@ pub fn decodeBlock(...@@ -675,7 +678,9 @@ pub fn decodeBlock(
675 bytes_written += len;678 bytes_written += len;
676 }679 }
677680
678 consumed_count.* += bytes_read;681 if (!decode_state.isLiteralStreamEmpty()) return error.MalformedCompressedBlock;
682
683 consumed_count.* += block_size;
679 return bytes_written;684 return bytes_written;
680 },685 },
681 .reserved => return error.ReservedBlock,686 .reserved => return error.ReservedBlock,
...@@ -749,13 +754,10 @@ pub fn decodeBlockRingBuffer(...@@ -749,13 +754,10 @@ pub fn decodeBlockRingBuffer(
749 sequence_size_limit -= decompressed_size;754 sequence_size_limit -= decompressed_size;
750 }755 }
751756
752 if (bit_stream.bit_reader.bit_count != 0) {757 if (!bit_stream.isEmpty()) {
753 return error.MalformedCompressedBlock;758 return error.MalformedCompressedBlock;
754 }759 }
755
756 bytes_read += bit_stream_bytes.len;
757 }760 }
758 if (bytes_read != block_size) return error.MalformedCompressedBlock;
759761
760 if (decode_state.literal_written_count < literals.header.regenerated_size) {762 if (decode_state.literal_written_count < literals.header.regenerated_size) {
761 const len = literals.header.regenerated_size - decode_state.literal_written_count;763 const len = literals.header.regenerated_size - decode_state.literal_written_count;
...@@ -764,7 +766,9 @@ pub fn decodeBlockRingBuffer(...@@ -764,7 +766,9 @@ pub fn decodeBlockRingBuffer(
764 bytes_written += len;766 bytes_written += len;
765 }767 }
766768
767 consumed_count.* += bytes_read;769 if (!decode_state.isLiteralStreamEmpty()) return error.MalformedCompressedBlock;
770
771 consumed_count.* += block_size;
768 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;772 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
769 return bytes_written;773 return bytes_written;
770 },774 },
...@@ -837,7 +841,7 @@ pub fn decodeBlockReader(...@@ -837,7 +841,7 @@ pub fn decodeBlockReader(
837 sequence_size_limit -= decompressed_size;841 sequence_size_limit -= decompressed_size;
838 bytes_written += decompressed_size;842 bytes_written += decompressed_size;
839 }843 }
840 if (bit_stream.bit_reader.bit_count != 0) {844 if (!bit_stream.isEmpty()) {
841 return error.MalformedCompressedBlock;845 return error.MalformedCompressedBlock;
842 }846 }
843 }847 }
...@@ -849,6 +853,8 @@ pub fn decodeBlockReader(...@@ -849,6 +853,8 @@ pub fn decodeBlockReader(
849 bytes_written += len;853 bytes_written += len;
850 }854 }
851855
856 if (!decode_state.isLiteralStreamEmpty()) return error.MalformedCompressedBlock;
857
852 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;858 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
853 if (block_reader_limited.bytes_left != 0) return error.MalformedCompressedBlock;859 if (block_reader_limited.bytes_left != 0) return error.MalformedCompressedBlock;
854 decode_state.literal_written_count = 0;860 decode_state.literal_written_count = 0;
lib/std/compress/zstandard/decode/huffman.zig+4
...@@ -86,6 +86,10 @@ fn assignWeights(huff_bits: *readers.ReverseBitReader, accuracy_log: usize, entr...@@ -86,6 +86,10 @@ fn assignWeights(huff_bits: *readers.ReverseBitReader, accuracy_log: usize, entr
86 odd_state = odd_data.baseline + odd_bits;86 odd_state = odd_data.baseline + odd_bits;
87 } else return error.MalformedHuffmanTree;87 } else return error.MalformedHuffmanTree;
8888
89 if (!huff_bits.isEmpty()) {
90 return error.MalformedHuffmanTree;
91 }
92
89 return i + 1; // stream contains all but the last symbol93 return i + 1; // stream contains all but the last symbol
90}94}
9195
lib/std/compress/zstandard/readers.zig+7-1
...@@ -36,7 +36,9 @@ pub const ReverseBitReader = struct {...@@ -36,7 +36,9 @@ pub const ReverseBitReader = struct {
36 pub fn init(self: *ReverseBitReader, bytes: []const u8) error{BitStreamHasNoStartBit}!void {36 pub fn init(self: *ReverseBitReader, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
37 self.byte_reader = ReversedByteReader.init(bytes);37 self.byte_reader = ReversedByteReader.init(bytes);
38 self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader());38 self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader());
39 while (0 == self.readBitsNoEof(u1, 1) catch return error.BitStreamHasNoStartBit) {}39 var i: usize = 0;
40 while (i < 8 and 0 == self.readBitsNoEof(u1, 1) catch return error.BitStreamHasNoStartBit) : (i += 1) {}
41 if (i == 8) return error.BitStreamHasNoStartBit;
40 }42 }
4143
42 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) error{EndOfStream}!U {44 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) error{EndOfStream}!U {
...@@ -50,6 +52,10 @@ pub const ReverseBitReader = struct {...@@ -50,6 +52,10 @@ pub const ReverseBitReader = struct {
50 pub fn alignToByte(self: *@This()) void {52 pub fn alignToByte(self: *@This()) void {
51 self.bit_reader.alignToByte();53 self.bit_reader.alignToByte();
52 }54 }
55
56 pub fn isEmpty(self: ReverseBitReader) bool {
57 return self.byte_reader.remaining_bytes == 0 and self.bit_reader.bit_count == 0;
58 }
53};59};
5460
55pub fn BitReader(comptime Reader: type) type {61pub fn BitReader(comptime Reader: type) type {