authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-01-24 13:07:58+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log6b85373875e2a7a8ef9d74e20e8e827dc622a29c
treedcc3a338c44af76e06fa1adeab6bdeb4380682fa
parent082acd7f17358ed3a7787f284b48b754fefd8187

std.compress.zstandard: validate sequence lengths


1 files changed, 26 insertions(+), 12 deletions(-)

lib/std/compress/zstandard/decompress.zig+26-12
...@@ -271,10 +271,9 @@ pub const DecodeState = struct {...@@ -271,10 +271,9 @@ pub const DecodeState = struct {
271 literals: LiteralsSection,271 literals: LiteralsSection,
272 sequence: Sequence,272 sequence: Sequence,
273 ) !void {273 ) !void {
274 try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length);274 if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence;
275275
276 // TODO: should we validate offset against max_window_size?276 try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length);
277 assert(sequence.offset <= write_pos + sequence.literal_length);
278 const copy_start = write_pos + sequence.literal_length - sequence.offset;277 const copy_start = write_pos + sequence.literal_length - sequence.offset;
279 const copy_end = copy_start + sequence.match_length;278 const copy_end = copy_start + sequence.match_length;
280 // NOTE: we ignore the usage message for std.mem.copy and copy with dest.ptr >= src.ptr279 // NOTE: we ignore the usage message for std.mem.copy and copy with dest.ptr >= src.ptr
...@@ -288,8 +287,9 @@ pub const DecodeState = struct {...@@ -288,8 +287,9 @@ pub const DecodeState = struct {
288 literals: LiteralsSection,287 literals: LiteralsSection,
289 sequence: Sequence,288 sequence: Sequence,
290 ) !void {289 ) !void {
290 if (sequence.offset > dest.data.len) return error.MalformedSequence;
291
291 try self.decodeLiteralsRingBuffer(dest, literals, sequence.literal_length);292 try self.decodeLiteralsRingBuffer(dest, literals, sequence.literal_length);
292 // TODO: check that ring buffer window is full enough for match copies
293 const copy_slice = dest.sliceAt(dest.write_index + dest.data.len - sequence.offset, sequence.match_length);293 const copy_slice = dest.sliceAt(dest.write_index + dest.data.len - sequence.offset, sequence.match_length);
294 // TODO: would std.mem.copy and figuring out dest slice be better/faster?294 // TODO: would std.mem.copy and figuring out dest slice be better/faster?
295 for (copy_slice.first) |b| dest.writeAssumeCapacity(b);295 for (copy_slice.first) |b| dest.writeAssumeCapacity(b);
...@@ -302,9 +302,13 @@ pub const DecodeState = struct {...@@ -302,9 +302,13 @@ pub const DecodeState = struct {
302 write_pos: usize,302 write_pos: usize,
303 literals: LiteralsSection,303 literals: LiteralsSection,
304 bit_reader: anytype,304 bit_reader: anytype,
305 sequence_size_limit: usize,
305 last_sequence: bool,306 last_sequence: bool,
306 ) !usize {307 ) !usize {
307 const sequence = try self.nextSequence(bit_reader);308 const sequence = try self.nextSequence(bit_reader);
309 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
310 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
311
308 try self.executeSequenceSlice(dest, write_pos, literals, sequence);312 try self.executeSequenceSlice(dest, write_pos, literals, sequence);
309 log.debug("sequence decompressed into '{x}'", .{313 log.debug("sequence decompressed into '{x}'", .{
310 std.fmt.fmtSliceHexUpper(dest[write_pos .. write_pos + sequence.literal_length + sequence.match_length]),314 std.fmt.fmtSliceHexUpper(dest[write_pos .. write_pos + sequence.literal_length + sequence.match_length]),
...@@ -314,7 +318,7 @@ pub const DecodeState = struct {...@@ -314,7 +318,7 @@ pub const DecodeState = struct {
314 try self.updateState(.match, bit_reader);318 try self.updateState(.match, bit_reader);
315 try self.updateState(.offset, bit_reader);319 try self.updateState(.offset, bit_reader);
316 }320 }
317 return sequence.match_length + sequence.literal_length;321 return sequence_length;
318 }322 }
319323
320 pub fn decodeSequenceRingBuffer(324 pub fn decodeSequenceRingBuffer(
...@@ -322,12 +326,15 @@ pub const DecodeState = struct {...@@ -322,12 +326,15 @@ pub const DecodeState = struct {
322 dest: *RingBuffer,326 dest: *RingBuffer,
323 literals: LiteralsSection,327 literals: LiteralsSection,
324 bit_reader: anytype,328 bit_reader: anytype,
329 sequence_size_limit: usize,
325 last_sequence: bool,330 last_sequence: bool,
326 ) !usize {331 ) !usize {
327 const sequence = try self.nextSequence(bit_reader);332 const sequence = try self.nextSequence(bit_reader);
333 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
334 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
335
328 try self.executeSequenceRingBuffer(dest, literals, sequence);336 try self.executeSequenceRingBuffer(dest, literals, sequence);
329 if (std.options.log_level == .debug) {337 if (std.options.log_level == .debug) {
330 const sequence_length = sequence.literal_length + sequence.match_length;
331 const written_slice = dest.sliceLast(sequence_length);338 const written_slice = dest.sliceLast(sequence_length);
332 log.debug("sequence decompressed into '{x}{x}'", .{339 log.debug("sequence decompressed into '{x}{x}'", .{
333 std.fmt.fmtSliceHexUpper(written_slice.first),340 std.fmt.fmtSliceHexUpper(written_slice.first),
...@@ -339,7 +346,7 @@ pub const DecodeState = struct {...@@ -339,7 +346,7 @@ pub const DecodeState = struct {
339 try self.updateState(.match, bit_reader);346 try self.updateState(.match, bit_reader);
340 try self.updateState(.offset, bit_reader);347 try self.updateState(.offset, bit_reader);
341 }348 }
342 return sequence.match_length + sequence.literal_length;349 return sequence_length;
343 }350 }
344351
345 fn nextLiteralMultiStream(self: *DecodeState, literals: LiteralsSection) !void {352 fn nextLiteralMultiStream(self: *DecodeState, literals: LiteralsSection) !void {
...@@ -717,9 +724,9 @@ pub fn decodeBlock(...@@ -717,9 +724,9 @@ pub fn decodeBlock(
717 consumed_count: *usize,724 consumed_count: *usize,
718 written_count: usize,725 written_count: usize,
719) !usize {726) !usize {
720 const block_maximum_size = 1 << 17; // 128KiB727 const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB
721 const block_size = block_header.block_size;728 const block_size = block_header.block_size;
722 if (block_maximum_size < block_size) return error.BlockSizeOverMaximum;729 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
723 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)730 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)
724 switch (block_header.block_type) {731 switch (block_header.block_type) {
725 .raw => return decodeRawBlock(dest[written_count..], src, block_size, consumed_count),732 .raw => return decodeRawBlock(dest[written_count..], src, block_size, consumed_count),
...@@ -739,17 +746,21 @@ pub fn decodeBlock(...@@ -739,17 +746,21 @@ pub fn decodeBlock(
739746
740 try decode_state.readInitialFseState(&bit_stream);747 try decode_state.readInitialFseState(&bit_stream);
741748
749 var sequence_size_limit = block_size_max;
742 var i: usize = 0;750 var i: usize = 0;
743 while (i < sequences_header.sequence_count) : (i += 1) {751 while (i < sequences_header.sequence_count) : (i += 1) {
744 log.debug("decoding sequence {d}", .{i});752 log.debug("decoding sequence {d}", .{i});
753 const write_pos = written_count + bytes_written;
745 const decompressed_size = try decode_state.decodeSequenceSlice(754 const decompressed_size = try decode_state.decodeSequenceSlice(
746 dest,755 dest,
747 written_count + bytes_written,756 write_pos,
748 literals,757 literals,
749 &bit_stream,758 &bit_stream,
759 sequence_size_limit,
750 i == sequences_header.sequence_count - 1,760 i == sequences_header.sequence_count - 1,
751 );761 );
752 bytes_written += decompressed_size;762 bytes_written += decompressed_size;
763 sequence_size_limit -= decompressed_size;
753 }764 }
754765
755 bytes_read += bit_stream_bytes.len;766 bytes_read += bit_stream_bytes.len;
...@@ -781,10 +792,10 @@ pub fn decodeBlockRingBuffer(...@@ -781,10 +792,10 @@ pub fn decodeBlockRingBuffer(
781 block_header: frame.ZStandard.Block.Header,792 block_header: frame.ZStandard.Block.Header,
782 decode_state: *DecodeState,793 decode_state: *DecodeState,
783 consumed_count: *usize,794 consumed_count: *usize,
784 block_size_maximum: usize,795 block_size_max: usize,
785) !usize {796) !usize {
786 const block_size = block_header.block_size;797 const block_size = block_header.block_size;
787 if (block_size_maximum < block_size) return error.BlockSizeOverMaximum;798 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
788 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)799 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)
789 switch (block_header.block_type) {800 switch (block_header.block_type) {
790 .raw => return decodeRawBlockRingBuffer(dest, src, block_size, consumed_count),801 .raw => return decodeRawBlockRingBuffer(dest, src, block_size, consumed_count),
...@@ -804,6 +815,7 @@ pub fn decodeBlockRingBuffer(...@@ -804,6 +815,7 @@ pub fn decodeBlockRingBuffer(
804815
805 try decode_state.readInitialFseState(&bit_stream);816 try decode_state.readInitialFseState(&bit_stream);
806817
818 var sequence_size_limit = block_size_max;
807 var i: usize = 0;819 var i: usize = 0;
808 while (i < sequences_header.sequence_count) : (i += 1) {820 while (i < sequences_header.sequence_count) : (i += 1) {
809 log.debug("decoding sequence {d}", .{i});821 log.debug("decoding sequence {d}", .{i});
...@@ -811,9 +823,11 @@ pub fn decodeBlockRingBuffer(...@@ -811,9 +823,11 @@ pub fn decodeBlockRingBuffer(
811 dest,823 dest,
812 literals,824 literals,
813 &bit_stream,825 &bit_stream,
826 sequence_size_limit,
814 i == sequences_header.sequence_count - 1,827 i == sequences_header.sequence_count - 1,
815 );828 );
816 bytes_written += decompressed_size;829 bytes_written += decompressed_size;
830 sequence_size_limit -= decompressed_size;
817 }831 }
818832
819 bytes_read += bit_stream_bytes.len;833 bytes_read += bit_stream_bytes.len;