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 {
391391 try self.literal_stream_reader.init(bytes);
392392 }
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
394401 const LiteralBitsError = error{
395402 BitStreamHasNoStartBit,
396403 UnexpectedEndOfLiteralStream,
397404 };
398405 fn readLiteralsBits(
399406 self: *DecodeState,
400 comptime T: type,
401407 bit_count_to_read: usize,
402 ) LiteralBitsError!T {
408 ) LiteralBitsError!u16 {
403409 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {
404410 if (self.literal_streams == .four and self.literal_stream_index < 3) {
405411 try self.nextLiteralMultiStream();
......@@ -461,7 +467,7 @@ pub const DecodeState = struct {
461467 while (i < len) : (i += 1) {
462468 var prefix: u16 = 0;
463469 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| {
465471 return err;
466472 };
467473 prefix <<= bit_count_to_read;
......@@ -533,7 +539,7 @@ pub const DecodeState = struct {
533539 while (i < len) : (i += 1) {
534540 var prefix: u16 = 0;
535541 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);
537543 prefix <<= bit_count_to_read;
538544 prefix |= new_bits;
539545 bits_read += bit_count_to_read;
......@@ -659,13 +665,10 @@ pub fn decodeBlock(
659665 sequence_size_limit -= decompressed_size;
660666 }
661667
662 if (bit_stream.bit_reader.bit_count != 0) {
668 if (!bit_stream.isEmpty()) {
663669 return error.MalformedCompressedBlock;
664670 }
665
666 bytes_read += bit_stream_bytes.len;
667671 }
668 if (bytes_read != block_size) return error.MalformedCompressedBlock;
669672
670673 if (decode_state.literal_written_count < literals.header.regenerated_size) {
671674 const len = literals.header.regenerated_size - decode_state.literal_written_count;
......@@ -675,7 +678,9 @@ pub fn decodeBlock(
675678 bytes_written += len;
676679 }
677680
678 consumed_count.* += bytes_read;
681 if (!decode_state.isLiteralStreamEmpty()) return error.MalformedCompressedBlock;
682
683 consumed_count.* += block_size;
679684 return bytes_written;
680685 },
681686 .reserved => return error.ReservedBlock,
......@@ -749,13 +754,10 @@ pub fn decodeBlockRingBuffer(
749754 sequence_size_limit -= decompressed_size;
750755 }
751756
752 if (bit_stream.bit_reader.bit_count != 0) {
757 if (!bit_stream.isEmpty()) {
753758 return error.MalformedCompressedBlock;
754759 }
755
756 bytes_read += bit_stream_bytes.len;
757760 }
758 if (bytes_read != block_size) return error.MalformedCompressedBlock;
759761
760762 if (decode_state.literal_written_count < literals.header.regenerated_size) {
761763 const len = literals.header.regenerated_size - decode_state.literal_written_count;
......@@ -764,7 +766,9 @@ pub fn decodeBlockRingBuffer(
764766 bytes_written += len;
765767 }
766768
767 consumed_count.* += bytes_read;
769 if (!decode_state.isLiteralStreamEmpty()) return error.MalformedCompressedBlock;
770
771 consumed_count.* += block_size;
768772 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
769773 return bytes_written;
770774 },
......@@ -837,7 +841,7 @@ pub fn decodeBlockReader(
837841 sequence_size_limit -= decompressed_size;
838842 bytes_written += decompressed_size;
839843 }
840 if (bit_stream.bit_reader.bit_count != 0) {
844 if (!bit_stream.isEmpty()) {
841845 return error.MalformedCompressedBlock;
842846 }
843847 }
......@@ -849,6 +853,8 @@ pub fn decodeBlockReader(
849853 bytes_written += len;
850854 }
851855
856 if (!decode_state.isLiteralStreamEmpty()) return error.MalformedCompressedBlock;
857
852858 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
853859 if (block_reader_limited.bytes_left != 0) return error.MalformedCompressedBlock;
854860 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
8686 odd_state = odd_data.baseline + odd_bits;
8787 } else return error.MalformedHuffmanTree;
8888
89 if (!huff_bits.isEmpty()) {
90 return error.MalformedHuffmanTree;
91 }
92
8993 return i + 1; // stream contains all but the last symbol
9094}
9195
lib/std/compress/zstandard/readers.zig+7-1
......@@ -36,7 +36,9 @@ pub const ReverseBitReader = struct {
3636 pub fn init(self: *ReverseBitReader, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
3737 self.byte_reader = ReversedByteReader.init(bytes);
3838 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;
4042 }
4143
4244 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) error{EndOfStream}!U {
......@@ -50,6 +52,10 @@ pub const ReverseBitReader = struct {
5052 pub fn alignToByte(self: *@This()) void {
5153 self.bit_reader.alignToByte();
5254 }
55
56 pub fn isEmpty(self: ReverseBitReader) bool {
57 return self.byte_reader.remaining_bytes == 0 and self.bit_reader.bit_count == 0;
58 }
5359};
5460
5561pub fn BitReader(comptime Reader: type) type {