authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-06 13:19:24+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log98bbd959b08db4d66a9ae68ce3fe4196594c6d9e
treeb5cd9926fba54cbfbb4b22d89eaafab14dc0a72a
parentece52e0771560726a72a11cfafb59bb1dc9ad221

std.compress.zstandard: improve block size validation


2 files changed, 8 insertions(+), 2 deletions(-)

lib/std/compress/zstandard/decode/block.zig+5-1
...@@ -594,9 +594,9 @@ pub fn decodeBlock(...@@ -594,9 +594,9 @@ pub fn decodeBlock(
594 block_header: frame.Zstandard.Block.Header,594 block_header: frame.Zstandard.Block.Header,
595 decode_state: *DecodeState,595 decode_state: *DecodeState,
596 consumed_count: *usize,596 consumed_count: *usize,
597 block_size_max: usize,
597 written_count: usize,598 written_count: usize,
598) (error{DestTooSmall} || Error)!usize {599) (error{DestTooSmall} || Error)!usize {
599 const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB
600 const block_size = block_header.block_size;600 const block_size = block_header.block_size;
601 if (block_size_max < block_size) return error.BlockSizeOverMaximum;601 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
602 switch (block_header.block_type) {602 switch (block_header.block_type) {
...@@ -805,6 +805,7 @@ pub fn decodeBlockReader(...@@ -805,6 +805,7 @@ pub fn decodeBlockReader(
805805
806 try decode_state.prepare(block_reader, literals, sequences_header);806 try decode_state.prepare(block_reader, literals, sequences_header);
807807
808 var bytes_written: usize = 0;
808 if (sequences_header.sequence_count > 0) {809 if (sequences_header.sequence_count > 0) {
809 if (sequence_buffer.len < block_reader_limited.bytes_left)810 if (sequence_buffer.len < block_reader_limited.bytes_left)
810 return error.SequenceBufferTooSmall;811 return error.SequenceBufferTooSmall;
...@@ -825,6 +826,7 @@ pub fn decodeBlockReader(...@@ -825,6 +826,7 @@ pub fn decodeBlockReader(
825 i == sequences_header.sequence_count - 1,826 i == sequences_header.sequence_count - 1,
826 ) catch return error.MalformedCompressedBlock;827 ) catch return error.MalformedCompressedBlock;
827 sequence_size_limit -= decompressed_size;828 sequence_size_limit -= decompressed_size;
829 bytes_written += decompressed_size;
828 }830 }
829 }831 }
830832
...@@ -832,8 +834,10 @@ pub fn decodeBlockReader(...@@ -832,8 +834,10 @@ pub fn decodeBlockReader(
832 const len = literals.header.regenerated_size - decode_state.literal_written_count;834 const len = literals.header.regenerated_size - decode_state.literal_written_count;
833 decode_state.decodeLiteralsRingBuffer(dest, len) catch835 decode_state.decodeLiteralsRingBuffer(dest, len) catch
834 return error.MalformedCompressedBlock;836 return error.MalformedCompressedBlock;
837 bytes_written += len;
835 }838 }
836839
840 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
837 decode_state.literal_written_count = 0;841 decode_state.literal_written_count = 0;
838 assert(block_reader.readByte() == error.EndOfStream);842 assert(block_reader.readByte() == error.EndOfStream);
839 },843 },
lib/std/compress/zstandard/decompress.zig+3-1
...@@ -236,6 +236,7 @@ pub fn decodeZstandardFrame(...@@ -236,6 +236,7 @@ pub fn decodeZstandardFrame(
236 src[consumed_count..],236 src[consumed_count..],
237 &consumed_count,237 &consumed_count,
238 if (frame_context.hasher_opt) |*hasher| hasher else null,238 if (frame_context.hasher_opt) |*hasher| hasher else null,
239 frame_context.block_size_max,
239 ) catch |err| switch (err) {240 ) catch |err| switch (err) {
240 error.DestTooSmall => return error.BadContentSize,241 error.DestTooSmall => return error.BadContentSize,
241 inline else => |e| return e,242 inline else => |e| return e,
...@@ -376,7 +377,6 @@ pub fn decodeZstandardFrameArrayList(...@@ -376,7 +377,6 @@ pub fn decodeZstandardFrameArrayList(
376 block_header = try block.decodeBlockHeaderSlice(src[consumed_count..]);377 block_header = try block.decodeBlockHeaderSlice(src[consumed_count..]);
377 consumed_count += 3;378 consumed_count += 3;
378 }) {379 }) {
379 if (block_header.block_size > frame_context.block_size_max) return error.BlockSizeOverMaximum;
380 const written_size = try block.decodeBlockRingBuffer(380 const written_size = try block.decodeBlockRingBuffer(
381 &ring_buffer,381 &ring_buffer,
382 src[consumed_count..],382 src[consumed_count..],
...@@ -420,6 +420,7 @@ fn decodeFrameBlocks(...@@ -420,6 +420,7 @@ fn decodeFrameBlocks(
420 src: []const u8,420 src: []const u8,
421 consumed_count: *usize,421 consumed_count: *usize,
422 hash: ?*std.hash.XxHash64,422 hash: ?*std.hash.XxHash64,
423 block_size_max: usize,
423) (error{ EndOfStream, DestTooSmall } || block.Error)!usize {424) (error{ EndOfStream, DestTooSmall } || block.Error)!usize {
424 // These tables take 7680 bytes425 // These tables take 7680 bytes
425 var literal_fse_data: [types.compressed_block.table_size_max.literal]Table.Fse = undefined;426 var literal_fse_data: [types.compressed_block.table_size_max.literal]Table.Fse = undefined;
...@@ -441,6 +442,7 @@ fn decodeFrameBlocks(...@@ -441,6 +442,7 @@ fn decodeFrameBlocks(
441 block_header,442 block_header,
442 &decode_state,443 &decode_state,
443 &bytes_read,444 &bytes_read,
445 block_size_max,
444 written_count,446 written_count,
445 );447 );
446 if (hash) |hash_state| hash_state.update(dest[written_count .. written_count + written_size]);448 if (hash) |hash_state| hash_state.update(dest[written_count .. written_count + written_size]);