authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-01-24 14:30:32+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log774e2f5a5c918cccfc455bcb73d90be43ec9a9eb
treea3dd2b546883fae1191860c61e17065b863d472a
parent31d1cae8c68fbc765fd4394863b071788dbc9746

std.compress.zstandard: add input length safety checks


1 files changed, 37 insertions(+), 19 deletions(-)

lib/std/compress/zstandard/decompress.zig+37-19
...@@ -680,7 +680,8 @@ pub fn decodeFrameBlocks(dest: []u8, src: []const u8, consumed_count: *usize, ha...@@ -680,7 +680,8 @@ pub fn decodeFrameBlocks(dest: []u8, src: []const u8, consumed_count: *usize, ha
680 return written_count;680 return written_count;
681}681}
682682
683fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) usize {683fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
684 if (src.len < block_size) return error.MalformedBlockSize;
684 log.debug("writing raw block - size {d}", .{block_size});685 log.debug("writing raw block - size {d}", .{block_size});
685 const data = src[0..block_size];686 const data = src[0..block_size];
686 std.mem.copy(u8, dest, data);687 std.mem.copy(u8, dest, data);
...@@ -688,7 +689,8 @@ fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:...@@ -688,7 +689,8 @@ fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:
688 return block_size;689 return block_size;
689}690}
690691
691fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) usize {692fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
693 if (src.len < block_size) return error.MalformedBlockSize;
692 log.debug("writing raw block - size {d}", .{block_size});694 log.debug("writing raw block - size {d}", .{block_size});
693 const data = src[0..block_size];695 const data = src[0..block_size];
694 dest.writeSliceAssumeCapacity(data);696 dest.writeSliceAssumeCapacity(data);
...@@ -696,7 +698,8 @@ fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21,...@@ -696,7 +698,8 @@ fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21,
696 return block_size;698 return block_size;
697}699}
698700
699fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) usize {701fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
702 if (src.len < 1) return error.MalformedRleBlock;
700 log.debug("writing rle block - '{x}'x{d}", .{ src[0], block_size });703 log.debug("writing rle block - '{x}'x{d}", .{ src[0], block_size });
701 var write_pos: usize = 0;704 var write_pos: usize = 0;
702 while (write_pos < block_size) : (write_pos += 1) {705 while (write_pos < block_size) : (write_pos += 1) {
...@@ -706,7 +709,8 @@ fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:...@@ -706,7 +709,8 @@ fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:
706 return block_size;709 return block_size;
707}710}
708711
709fn decodeRleBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) usize {712fn decodeRleBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
713 if (src.len < 1) return error.MalformedRleBlock;
710 log.debug("writing rle block - '{x}'x{d}", .{ src[0], block_size });714 log.debug("writing rle block - '{x}'x{d}", .{ src[0], block_size });
711 var write_pos: usize = 0;715 var write_pos: usize = 0;
712 while (write_pos < block_size) : (write_pos += 1) {716 while (write_pos < block_size) : (write_pos += 1) {
...@@ -727,11 +731,11 @@ pub fn decodeBlock(...@@ -727,11 +731,11 @@ pub fn decodeBlock(
727 const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB731 const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB
728 const block_size = block_header.block_size;732 const block_size = block_header.block_size;
729 if (block_size_max < block_size) return error.BlockSizeOverMaximum;733 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
730 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)
731 switch (block_header.block_type) {734 switch (block_header.block_type) {
732 .raw => return decodeRawBlock(dest[written_count..], src, block_size, consumed_count),735 .raw => return decodeRawBlock(dest[written_count..], src, block_size, consumed_count),
733 .rle => return decodeRleBlock(dest[written_count..], src, block_size, consumed_count),736 .rle => return decodeRleBlock(dest[written_count..], src, block_size, consumed_count),
734 .compressed => {737 .compressed => {
738 if (src.len < block_size) return error.MalformedBlockSize;
735 var bytes_read: usize = 0;739 var bytes_read: usize = 0;
736 const literals = try decodeLiteralsSection(src, &bytes_read);740 const literals = try decodeLiteralsSection(src, &bytes_read);
737 const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read);741 const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read);
...@@ -796,11 +800,11 @@ pub fn decodeBlockRingBuffer(...@@ -796,11 +800,11 @@ pub fn decodeBlockRingBuffer(
796) !usize {800) !usize {
797 const block_size = block_header.block_size;801 const block_size = block_header.block_size;
798 if (block_size_max < block_size) return error.BlockSizeOverMaximum;802 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
799 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)
800 switch (block_header.block_type) {803 switch (block_header.block_type) {
801 .raw => return decodeRawBlockRingBuffer(dest, src, block_size, consumed_count),804 .raw => return decodeRawBlockRingBuffer(dest, src, block_size, consumed_count),
802 .rle => return decodeRleBlockRingBuffer(dest, src, block_size, consumed_count),805 .rle => return decodeRleBlockRingBuffer(dest, src, block_size, consumed_count),
803 .compressed => {806 .compressed => {
807 if (src.len < block_size) return error.MalformedBlockSize;
804 var bytes_read: usize = 0;808 var bytes_read: usize = 0;
805 const literals = try decodeLiteralsSection(src, &bytes_read);809 const literals = try decodeLiteralsSection(src, &bytes_read);
806 const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read);810 const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read);
...@@ -957,11 +961,11 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {...@@ -957,11 +961,11 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {
957}961}
958962
959pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsSection {963pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsSection {
960 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)
961 var bytes_read: usize = 0;964 var bytes_read: usize = 0;
962 const header = decodeLiteralsHeader(src, &bytes_read);965 const header = try decodeLiteralsHeader(src, &bytes_read);
963 switch (header.block_type) {966 switch (header.block_type) {
964 .raw => {967 .raw => {
968 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;
965 const stream = src[bytes_read .. bytes_read + header.regenerated_size];969 const stream = src[bytes_read .. bytes_read + header.regenerated_size];
966 consumed_count.* += header.regenerated_size + bytes_read;970 consumed_count.* += header.regenerated_size + bytes_read;
967 return LiteralsSection{971 return LiteralsSection{
...@@ -971,6 +975,7 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS...@@ -971,6 +975,7 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS
971 };975 };
972 },976 },
973 .rle => {977 .rle => {
978 if (src.len < bytes_read + 1) return error.MalformedLiteralsSection;
974 const stream = src[bytes_read .. bytes_read + 1];979 const stream = src[bytes_read .. bytes_read + 1];
975 consumed_count.* += 1 + bytes_read;980 consumed_count.* += 1 + bytes_read;
976 return LiteralsSection{981 return LiteralsSection{
...@@ -990,18 +995,19 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS...@@ -990,18 +995,19 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS
990 log.debug("huffman tree size = {}, total streams size = {}", .{ huffman_tree_size, total_streams_size });995 log.debug("huffman tree size = {}, total streams size = {}", .{ huffman_tree_size, total_streams_size });
991 if (huffman_tree) |tree| dumpHuffmanTree(tree);996 if (huffman_tree) |tree| dumpHuffmanTree(tree);
992997
998 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;
999 const stream_data = src[bytes_read .. bytes_read + total_streams_size];
1000
993 if (header.size_format == 0) {1001 if (header.size_format == 0) {
994 const stream = src[bytes_read .. bytes_read + total_streams_size];1002 consumed_count.* += total_streams_size + bytes_read;
995 bytes_read += total_streams_size;
996 consumed_count.* += bytes_read;
997 return LiteralsSection{1003 return LiteralsSection{
998 .header = header,1004 .header = header,
999 .huffman_tree = huffman_tree,1005 .huffman_tree = huffman_tree,
1000 .streams = .{ .one = stream },1006 .streams = .{ .one = stream_data },
1001 };1007 };
1002 }1008 }
10031009
1004 const stream_data = src[bytes_read .. bytes_read + total_streams_size];1010 if (stream_data.len < 6) return error.MalformedLiteralsSection;
10051011
1006 log.debug("jump table: {}", .{std.fmt.fmtSliceHexUpper(stream_data[0..6])});1012 log.debug("jump table: {}", .{std.fmt.fmtSliceHexUpper(stream_data[0..6])});
1007 const stream_1_length = @as(usize, readInt(u16, stream_data[0..2]));1013 const stream_1_length = @as(usize, readInt(u16, stream_data[0..2]));
...@@ -1014,6 +1020,7 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS...@@ -1014,6 +1020,7 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS
1014 const stream_3_start = stream_2_start + stream_2_length;1020 const stream_3_start = stream_2_start + stream_2_length;
1015 const stream_4_start = stream_3_start + stream_3_length;1021 const stream_4_start = stream_3_start + stream_3_length;
10161022
1023 if (stream_data.len < stream_4_start + stream_4_length) return error.MalformedLiteralsSection;
1017 consumed_count.* += total_streams_size + bytes_read;1024 consumed_count.* += total_streams_size + bytes_read;
10181025
1019 return LiteralsSection{1026 return LiteralsSection{
...@@ -1033,13 +1040,15 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS...@@ -1033,13 +1040,15 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS
1033fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.HuffmanTree {1040fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.HuffmanTree {
1034 var bytes_read: usize = 0;1041 var bytes_read: usize = 0;
1035 bytes_read += 1;1042 bytes_read += 1;
1043 if (src.len == 0) return error.MalformedHuffmanTree;
1036 const header = src[0];1044 const header = src[0];
1037 var symbol_count: usize = undefined;1045 var symbol_count: usize = undefined;
1038 var weights: [256]u4 = undefined;1046 var weights: [256]u4 = undefined;
1039 var max_number_of_bits: u4 = undefined;1047 var max_number_of_bits: u4 = undefined;
1040 if (header < 128) {1048 if (header < 128) {
1041 // FSE compressed weigths1049 // FSE compressed weights
1042 const compressed_size = header;1050 const compressed_size = header;
1051 if (src.len < 1 + compressed_size) return error.MalformedHuffmanTree;
1043 var stream = std.io.fixedBufferStream(src[1 .. compressed_size + 1]);1052 var stream = std.io.fixedBufferStream(src[1 .. compressed_size + 1]);
1044 var counting_reader = std.io.countingReader(stream.reader());1053 var counting_reader = std.io.countingReader(stream.reader());
1045 var bit_reader = bitReader(counting_reader.reader());1054 var bit_reader = bitReader(counting_reader.reader());
...@@ -1185,8 +1194,8 @@ fn lessThanByWeight(...@@ -1185,8 +1194,8 @@ fn lessThanByWeight(
1185 return weights[lhs.symbol] < weights[rhs.symbol];1194 return weights[lhs.symbol] < weights[rhs.symbol];
1186}1195}
11871196
1188pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) LiteralsSection.Header {1197pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSection.Header {
1189 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)1198 if (src.len == 0) return error.MalformedLiteralsSection;
1190 const start = consumed_count.*;1199 const start = consumed_count.*;
1191 const byte0 = src[0];1200 const byte0 = src[0];
1192 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);1201 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);
...@@ -1201,14 +1210,16 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) LiteralsSec...@@ -1201,14 +1210,16 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) LiteralsSec
1201 consumed_count.* += 1;1210 consumed_count.* += 1;
1202 },1211 },
1203 1 => {1212 1 => {
1213 if (src.len < 2) return error.MalformedLiteralsHeader;
1204 regenerated_size = (byte0 >> 4) +1214 regenerated_size = (byte0 >> 4) +
1205 (@as(u20, src[consumed_count.* + 1]) << 4);1215 (@as(u20, src[1]) << 4);
1206 consumed_count.* += 2;1216 consumed_count.* += 2;
1207 },1217 },
1208 3 => {1218 3 => {
1219 if (src.len < 3) return error.MalformedLiteralsHeader;
1209 regenerated_size = (byte0 >> 4) +1220 regenerated_size = (byte0 >> 4) +
1210 (@as(u20, src[consumed_count.* + 1]) << 4) +1221 (@as(u20, src[1]) << 4) +
1211 (@as(u20, src[consumed_count.* + 2]) << 12);1222 (@as(u20, src[2]) << 12);
1212 consumed_count.* += 3;1223 consumed_count.* += 3;
1213 },1224 },
1214 }1225 }
...@@ -1218,17 +1229,20 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) LiteralsSec...@@ -1218,17 +1229,20 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) LiteralsSec
1218 const byte2 = src[2];1229 const byte2 = src[2];
1219 switch (size_format) {1230 switch (size_format) {
1220 0, 1 => {1231 0, 1 => {
1232 if (src.len < 3) return error.MalformedLiteralsHeader;
1221 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);1233 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);
1222 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);1234 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);
1223 consumed_count.* += 3;1235 consumed_count.* += 3;
1224 },1236 },
1225 2 => {1237 2 => {
1238 if (src.len < 4) return error.MalformedLiteralsHeader;
1226 const byte3 = src[3];1239 const byte3 = src[3];
1227 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);1240 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);
1228 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);1241 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);
1229 consumed_count.* += 4;1242 consumed_count.* += 4;
1230 },1243 },
1231 3 => {1244 3 => {
1245 if (src.len < 5) return error.MalformedLiteralsHeader;
1232 const byte3 = src[3];1246 const byte3 = src[3];
1233 const byte4 = src[4];1247 const byte4 = src[4];
1234 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);1248 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);
...@@ -1257,6 +1271,7 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) LiteralsSec...@@ -1257,6 +1271,7 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) LiteralsSec
1257}1271}
12581272
1259pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !SequencesSection.Header {1273pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !SequencesSection.Header {
1274 if (src.len == 0) return error.MalformedSequencesSection;
1260 var sequence_count: u24 = undefined;1275 var sequence_count: u24 = undefined;
12611276
1262 var bytes_read: usize = 0;1277 var bytes_read: usize = 0;
...@@ -1275,13 +1290,16 @@ pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !Sequences...@@ -1275,13 +1290,16 @@ pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !Sequences
1275 sequence_count = byte0;1290 sequence_count = byte0;
1276 bytes_read += 1;1291 bytes_read += 1;
1277 } else if (byte0 < 255) {1292 } else if (byte0 < 255) {
1293 if (src.len < 2) return error.MalformedSequencesSection;
1278 sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1];1294 sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1];
1279 bytes_read += 2;1295 bytes_read += 2;
1280 } else {1296 } else {
1297 if (src.len < 3) return error.MalformedSequencesSection;
1281 sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00;1298 sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00;
1282 bytes_read += 3;1299 bytes_read += 3;
1283 }1300 }
12841301
1302 if (src.len < bytes_read + 1) return error.MalformedSequencesSection;
1285 const compression_modes = src[bytes_read];1303 const compression_modes = src[bytes_read];
1286 bytes_read += 1;1304 bytes_read += 1;
12871305