authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-11-06 15:31:17+11:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2023-11-10 15:18:16-05:00
log138a35df8f434115be04641b1df29514b0ef1cb8
tree15ff4de3ba769e5c1a3b4b9f11202cbc157d005b
parent9ad03b628f5d4770f9f26e646019292b5ae9cf9b

zstandard: fix division by zero when using RingBuffer

This change fixes some division-by-zero bugs introduced by the optimized ring buffer read/write functions in d8c067966. There are edge cases where decompression can use a length zero ring buffer as the size of the ring buffer used is exactly the the window size specified by a Zstandard frame, and this can be zero. Switching away from loops to mem copies means that we need to ensure ring buffers do not have length zero ring when attempting to read/write from them.

2 files changed, 63 insertions(+), 14 deletions(-)

lib/std/compress/zstandard.zig+53-8
...@@ -70,13 +70,11 @@ pub fn DecompressStream(...@@ -70,13 +70,11 @@ pub fn DecompressStream(
70 self.state = .NewFrame;70 self.state = .NewFrame;
71 },71 },
72 .zstandard => |header| {72 .zstandard => |header| {
73 const frame_context = context: {73 const frame_context = try decompress.FrameContext.init(
74 break :context try decompress.FrameContext.init(74 header,
75 header,75 options.window_size_max,
76 options.window_size_max,76 options.verify_checksum,
77 options.verify_checksum,77 );
78 );
79 };
8078
81 const literal_fse_buffer = try self.allocator.alloc(79 const literal_fse_buffer = try self.allocator.alloc(
82 types.compressed_block.Table.Fse,80 types.compressed_block.Table.Fse,
...@@ -219,7 +217,9 @@ pub fn DecompressStream(...@@ -219,7 +217,9 @@ pub fn DecompressStream(
219 }217 }
220218
221 const size = @min(self.buffer.len(), buffer.len);219 const size = @min(self.buffer.len(), buffer.len);
222 self.buffer.readFirstAssumeLength(buffer, size);220 if (size > 0) {
221 self.buffer.readFirstAssumeLength(buffer, size);
222 }
223 if (self.state == .LastBlock and self.buffer.len() == 0) {223 if (self.state == .LastBlock and self.buffer.len() == 0) {
224 self.state = .NewFrame;224 self.state = .NewFrame;
225 self.allocator.free(self.literal_fse_buffer);225 self.allocator.free(self.literal_fse_buffer);
...@@ -282,3 +282,48 @@ test "zstandard decompression" {...@@ -282,3 +282,48 @@ test "zstandard decompression" {
282 try testReader(compressed3, uncompressed);282 try testReader(compressed3, uncompressed);
283 try testReader(compressed19, uncompressed);283 try testReader(compressed19, uncompressed);
284}284}
285
286fn expectEqualDecoded(expected: []const u8, input: []const u8) !void {
287 const allocator = std.testing.allocator;
288
289 {
290 const result = try decompress.decodeAlloc(allocator, input, false, 1 << 23);
291 defer allocator.free(result);
292 try std.testing.expectEqualStrings(expected, result);
293 }
294
295 {
296 var buffer = try allocator.alloc(u8, 2 * expected.len);
297 defer allocator.free(buffer);
298
299 const size = try decompress.decode(buffer, input, false);
300 try std.testing.expectEqualStrings(expected, buffer[0..size]);
301 }
302
303 {
304 var in_stream = std.io.fixedBufferStream(input);
305 var stream = decompressStream(allocator, in_stream.reader());
306 defer stream.deinit();
307
308 const result = try stream.reader().readAllAlloc(allocator, std.math.maxInt(usize));
309 defer allocator.free(result);
310
311 try std.testing.expectEqualStrings(expected, result);
312 }
313}
314
315test "zero sized block" {
316 const input_raw =
317 "\x28\xb5\x2f\xfd" ++ // zstandard frame magic number
318 "\x20\x00" ++ // frame header: only single_segment_flag set, frame_content_size zero
319 "\x01\x00\x00"; // block header with: last_block set, block_type raw, block_size zero
320
321 const input_rle =
322 "\x28\xb5\x2f\xfd" ++ // zstandard frame magic number
323 "\x20\x00" ++ // frame header: only single_segment_flag set, frame_content_size zero
324 "\x03\x00\x00" ++ // block header with: last_block set, block_type rle, block_size zero
325 "\xaa"; // block_content
326
327 try expectEqualDecoded("", input_raw);
328 try expectEqualDecoded("", input_rle);
329}
lib/std/compress/zstandard/decode/block.zig+10-6
...@@ -713,10 +713,14 @@ pub fn decodeBlockRingBuffer(...@@ -713,10 +713,14 @@ pub fn decodeBlockRingBuffer(
713 switch (block_header.block_type) {713 switch (block_header.block_type) {
714 .raw => {714 .raw => {
715 if (src.len < block_size) return error.MalformedBlockSize;715 if (src.len < block_size) return error.MalformedBlockSize;
716 const data = src[0..block_size];716 // dest may have length zero if block_size == 0, causing division by zero in
717 dest.writeSliceAssumeCapacity(data);717 // writeSliceAssumeCapacity()
718 consumed_count.* += block_size;718 if (block_size > 0) {
719 decode_state.written_count += block_size;719 const data = src[0..block_size];
720 dest.writeSliceAssumeCapacity(data);
721 consumed_count.* += block_size;
722 decode_state.written_count += block_size;
723 }
720 return block_size;724 return block_size;
721 },725 },
722 .rle => {726 .rle => {
...@@ -934,7 +938,7 @@ pub fn decodeLiteralsSectionSlice(...@@ -934,7 +938,7 @@ pub fn decodeLiteralsSectionSlice(
934 switch (header.block_type) {938 switch (header.block_type) {
935 .raw => {939 .raw => {
936 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;940 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;
937 const stream = src[bytes_read .. bytes_read + header.regenerated_size];941 const stream = src[bytes_read..][0..header.regenerated_size];
938 consumed_count.* += header.regenerated_size + bytes_read;942 consumed_count.* += header.regenerated_size + bytes_read;
939 return LiteralsSection{943 return LiteralsSection{
940 .header = header,944 .header = header,
...@@ -944,7 +948,7 @@ pub fn decodeLiteralsSectionSlice(...@@ -944,7 +948,7 @@ pub fn decodeLiteralsSectionSlice(
944 },948 },
945 .rle => {949 .rle => {
946 if (src.len < bytes_read + 1) return error.MalformedLiteralsSection;950 if (src.len < bytes_read + 1) return error.MalformedLiteralsSection;
947 const stream = src[bytes_read .. bytes_read + 1];951 const stream = src[bytes_read..][0..1];
948 consumed_count.* += 1 + bytes_read;952 consumed_count.* += 1 + bytes_read;
949 return LiteralsSection{953 return LiteralsSection{
950 .header = header,954 .header = header,