authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-03 12:56:51+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log596a97fb556a380d9f3780363ee07fc5dfaf43c2
tree3ce14f5683369683664fee3d5e4962987c493cb2
parenta651704876acfa8386c5eb117800f425e194bdc6

std.compress.zstandard: fix crashes


3 files changed, 33 insertions(+), 20 deletions(-)

lib/std/compress/zstandard/decode/block.zig+21-13
...@@ -148,8 +148,8 @@ pub const DecodeState = struct {...@@ -148,8 +148,8 @@ pub const DecodeState = struct {
148 }148 }
149149
150 fn updateRepeatOffset(self: *DecodeState, offset: u32) void {150 fn updateRepeatOffset(self: *DecodeState, offset: u32) void {
151 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[1]);151 self.repeat_offsets[2] = self.repeat_offsets[1];
152 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[2]);152 self.repeat_offsets[1] = self.repeat_offsets[0];
153 self.repeat_offsets[0] = offset;153 self.repeat_offsets[0] = offset;
154 }154 }
155155
...@@ -238,18 +238,22 @@ pub const DecodeState = struct {...@@ -238,18 +238,22 @@ pub const DecodeState = struct {
238 fn nextSequence(238 fn nextSequence(
239 self: *DecodeState,239 self: *DecodeState,
240 bit_reader: *readers.ReverseBitReader,240 bit_reader: *readers.ReverseBitReader,
241 ) error{ OffsetCodeTooLarge, EndOfStream }!Sequence {241 ) error{ InvalidBitStream, EndOfStream }!Sequence {
242 const raw_code = self.getCode(.offset);242 const raw_code = self.getCode(.offset);
243 const offset_code = std.math.cast(u5, raw_code) orelse {243 const offset_code = std.math.cast(u5, raw_code) orelse {
244 return error.OffsetCodeTooLarge;244 return error.InvalidBitStream;
245 };245 };
246 const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code);246 const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code);
247247
248 const match_code = self.getCode(.match);248 const match_code = self.getCode(.match);
249 if (match_code >= types.compressed_block.match_length_code_table.len)
250 return error.InvalidBitStream;
249 const match = types.compressed_block.match_length_code_table[match_code];251 const match = types.compressed_block.match_length_code_table[match_code];
250 const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]);252 const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]);
251253
252 const literal_code = self.getCode(.literal);254 const literal_code = self.getCode(.literal);
255 if (literal_code >= types.compressed_block.literals_length_code_table.len)
256 return error.InvalidBitStream;
253 const literal = types.compressed_block.literals_length_code_table[literal_code];257 const literal = types.compressed_block.literals_length_code_table[literal_code];
254 const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]);258 const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]);
255259
...@@ -269,6 +273,8 @@ pub const DecodeState = struct {...@@ -269,6 +273,8 @@ pub const DecodeState = struct {
269 break :offset self.useRepeatOffset(offset_value - 1);273 break :offset self.useRepeatOffset(offset_value - 1);
270 };274 };
271275
276 if (offset == 0) return error.InvalidBitStream;
277
272 return .{278 return .{
273 .literal_length = literal_length,279 .literal_length = literal_length,
274 .match_length = match_length,280 .match_length = match_length,
...@@ -308,7 +314,7 @@ pub const DecodeState = struct {...@@ -308,7 +314,7 @@ pub const DecodeState = struct {
308 }314 }
309315
310 const DecodeSequenceError = error{316 const DecodeSequenceError = error{
311 OffsetCodeTooLarge,317 InvalidBitStream,
312 EndOfStream,318 EndOfStream,
313 MalformedSequence,319 MalformedSequence,
314 MalformedFseBits,320 MalformedFseBits,
...@@ -326,7 +332,7 @@ pub const DecodeState = struct {...@@ -326,7 +332,7 @@ pub const DecodeState = struct {
326 /// - `error.UnexpectedEndOfLiteralStream` if the decoder state's literal332 /// - `error.UnexpectedEndOfLiteralStream` if the decoder state's literal
327 /// streams do not contain enough literals for the sequence (this may333 /// streams do not contain enough literals for the sequence (this may
328 /// mean the literal stream or the sequence is malformed).334 /// mean the literal stream or the sequence is malformed).
329 /// - `error.OffsetCodeTooLarge` if an invalid offset code is found335 /// - `error.InvalidBitStream` if the FSE sequence bitstream is malformed
330 /// - `error.EndOfStream` if `bit_reader` does not contain enough bits336 /// - `error.EndOfStream` if `bit_reader` does not contain enough bits
331 pub fn decodeSequenceSlice(337 pub fn decodeSequenceSlice(
332 self: *DecodeState,338 self: *DecodeState,
...@@ -608,9 +614,9 @@ pub fn decodeBlock(...@@ -608,9 +614,9 @@ pub fn decodeBlock(
608 .compressed => {614 .compressed => {
609 if (src.len < block_size) return error.MalformedBlockSize;615 if (src.len < block_size) return error.MalformedBlockSize;
610 var bytes_read: usize = 0;616 var bytes_read: usize = 0;
611 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch617 const literals = decodeLiteralsSectionSlice(src[0..block_size], &bytes_read) catch
612 return error.MalformedCompressedBlock;618 return error.MalformedCompressedBlock;
613 var fbs = std.io.fixedBufferStream(src[bytes_read..]);619 var fbs = std.io.fixedBufferStream(src[bytes_read..block_size]);
614 const fbs_reader = fbs.reader();620 const fbs_reader = fbs.reader();
615 const sequences_header = decodeSequencesHeader(fbs_reader) catch621 const sequences_header = decodeSequencesHeader(fbs_reader) catch
616 return error.MalformedCompressedBlock;622 return error.MalformedCompressedBlock;
...@@ -695,9 +701,9 @@ pub fn decodeBlockRingBuffer(...@@ -695,9 +701,9 @@ pub fn decodeBlockRingBuffer(
695 .compressed => {701 .compressed => {
696 if (src.len < block_size) return error.MalformedBlockSize;702 if (src.len < block_size) return error.MalformedBlockSize;
697 var bytes_read: usize = 0;703 var bytes_read: usize = 0;
698 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch704 const literals = decodeLiteralsSectionSlice(src[0..block_size], &bytes_read) catch
699 return error.MalformedCompressedBlock;705 return error.MalformedCompressedBlock;
700 var fbs = std.io.fixedBufferStream(src[bytes_read..]);706 var fbs = std.io.fixedBufferStream(src[bytes_read..block_size]);
701 const fbs_reader = fbs.reader();707 const fbs_reader = fbs.reader();
702 const sequences_header = decodeSequencesHeader(fbs_reader) catch708 const sequences_header = decodeSequencesHeader(fbs_reader) catch
703 return error.MalformedCompressedBlock;709 return error.MalformedCompressedBlock;
...@@ -894,7 +900,8 @@ pub fn decodeLiteralsSectionSlice(...@@ -894,7 +900,8 @@ pub fn decodeLiteralsSectionSlice(
894 else900 else
895 null;901 null;
896 const huffman_tree_size = bytes_read - huffman_tree_start;902 const huffman_tree_size = bytes_read - huffman_tree_start;
897 const total_streams_size = @as(usize, header.compressed_size.?) - huffman_tree_size;903 const total_streams_size = std.math.sub(usize, header.compressed_size.?, huffman_tree_size) catch
904 return error.MalformedLiteralsSection;
898905
899 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;906 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;
900 const stream_data = src[bytes_read .. bytes_read + total_streams_size];907 const stream_data = src[bytes_read .. bytes_read + total_streams_size];
...@@ -940,8 +947,9 @@ pub fn decodeLiteralsSection(...@@ -940,8 +947,9 @@ pub fn decodeLiteralsSection(
940 try huffman.decodeHuffmanTree(counting_reader.reader(), buffer)947 try huffman.decodeHuffmanTree(counting_reader.reader(), buffer)
941 else948 else
942 null;949 null;
943 const huffman_tree_size = counting_reader.bytes_read;950 const huffman_tree_size = @intCast(usize, counting_reader.bytes_read);
944 const total_streams_size = @as(usize, header.compressed_size.?) - @intCast(usize, huffman_tree_size);951 const total_streams_size = std.math.sub(usize, header.compressed_size.?, huffman_tree_size) catch
952 return error.MalformedLiteralsSection;
945953
946 if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall;954 if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall;
947 try source.readNoEof(buffer[0..total_streams_size]);955 try source.readNoEof(buffer[0..total_streams_size]);
lib/std/compress/zstandard/decode/huffman.zig+2-1
...@@ -146,13 +146,14 @@ fn assignSymbols(weight_sorted_prefixed_symbols: []LiteralsSection.HuffmanTree.P...@@ -146,13 +146,14 @@ fn assignSymbols(weight_sorted_prefixed_symbols: []LiteralsSection.HuffmanTree.P
146 return prefixed_symbol_count;146 return prefixed_symbol_count;
147}147}
148148
149fn buildHuffmanTree(weights: *[256]u4, symbol_count: usize) LiteralsSection.HuffmanTree {149fn buildHuffmanTree(weights: *[256]u4, symbol_count: usize) error{MalformedHuffmanTree}!LiteralsSection.HuffmanTree {
150 var weight_power_sum: u16 = 0;150 var weight_power_sum: u16 = 0;
151 for (weights[0 .. symbol_count - 1]) |value| {151 for (weights[0 .. symbol_count - 1]) |value| {
152 if (value > 0) {152 if (value > 0) {
153 weight_power_sum += @as(u16, 1) << (value - 1);153 weight_power_sum += @as(u16, 1) << (value - 1);
154 }154 }
155 }155 }
156 if (weight_power_sum >= 1 << 11) return error.MalformedHuffmanTree;
156157
157 // advance to next power of two (even if weight_power_sum is a power of 2)158 // advance to next power of two (even if weight_power_sum is a power of 2)
158 const max_number_of_bits = std.math.log2_int(u16, weight_power_sum) + 1;159 const max_number_of_bits = std.math.log2_int(u16, weight_power_sum) + 1;
lib/std/compress/zstandard/decompress.zig+10-6
...@@ -195,6 +195,7 @@ pub fn decodeZstandardFrame(...@@ -195,6 +195,7 @@ pub fn decodeZstandardFrame(
195 );195 );
196196
197 if (frame_header.descriptor.content_checksum_flag) {197 if (frame_header.descriptor.content_checksum_flag) {
198 if (src.len < consumed_count + 4) return error.EndOfStream;
198 const checksum = readIntSlice(u32, src[consumed_count .. consumed_count + 4]);199 const checksum = readIntSlice(u32, src[consumed_count .. consumed_count + 4]);
199 consumed_count += 4;200 consumed_count += 4;
200 if (hasher_opt) |*hasher| {201 if (hasher_opt) |*hasher| {
...@@ -302,17 +303,20 @@ pub fn decodeZstandardFrameAlloc(...@@ -302,17 +303,20 @@ pub fn decodeZstandardFrameAlloc(
302 &consumed_count,303 &consumed_count,
303 frame_context.block_size_max,304 frame_context.block_size_max,
304 );305 );
305 const written_slice = ring_buffer.sliceLast(written_size);306 if (written_size > 0) {
306 try result.appendSlice(written_slice.first);307 const written_slice = ring_buffer.sliceLast(written_size);
307 try result.appendSlice(written_slice.second);308 try result.appendSlice(written_slice.first);
308 if (frame_context.hasher_opt) |*hasher| {309 try result.appendSlice(written_slice.second);
309 hasher.update(written_slice.first);310 if (frame_context.hasher_opt) |*hasher| {
310 hasher.update(written_slice.second);311 hasher.update(written_slice.first);
312 hasher.update(written_slice.second);
313 }
311 }314 }
312 if (block_header.last_block) break;315 if (block_header.last_block) break;
313 }316 }
314317
315 if (frame_context.has_checksum) {318 if (frame_context.has_checksum) {
319 if (src.len < consumed_count + 4) return error.EndOfStream;
316 const checksum = readIntSlice(u32, src[consumed_count .. consumed_count + 4]);320 const checksum = readIntSlice(u32, src[consumed_count .. consumed_count + 4]);
317 consumed_count += 4;321 consumed_count += 4;
318 if (frame_context.hasher_opt) |*hasher| {322 if (frame_context.hasher_opt) |*hasher| {