authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-01-28 21:03:55+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log3bfba365483ccf30b197195cce8d5656f2c73736
treef032d254f267ef44ef71c43aba1ebcae2731a7de
parent3c06e2e7d0de92c0674c16ae23e1462a3acbe718

std.compress.zstandard: clean up error sets and line lengths


2 files changed, 215 insertions(+), 94 deletions(-)

lib/std/compress/zstandard/decompress.zig+214-93
......@@ -22,7 +22,7 @@ fn isSkippableMagic(magic: u32) bool {
2222/// if the the frame is skippable, `null` for Zstanndard frames that do not
2323/// declare their content size. Returns `UnusedBitSet` and `ReservedBitSet`
2424/// errors if the respective bits of the the frame descriptor are set.
25pub fn getFrameDecompressedSize(src: []const u8) !?u64 {
25pub fn getFrameDecompressedSize(src: []const u8) (InvalidBit || error{BadMagic})!?u64 {
2626 switch (try frameType(src)) {
2727 .zstandard => {
2828 const header = try decodeZStandardHeader(src[4..], null);
......@@ -52,7 +52,11 @@ const ReadWriteCount = struct {
5252
5353/// Decodes the frame at the start of `src` into `dest`. Returns the number of
5454/// bytes read from `src` and written to `dest`.
55pub fn decodeFrame(dest: []u8, src: []const u8, verify_checksum: bool) !ReadWriteCount {
55pub fn decodeFrame(
56 dest: []u8,
57 src: []const u8,
58 verify_checksum: bool,
59) (error{ UnknownContentSizeUnsupported, ContentTooLarge, BadMagic } || FrameError)!ReadWriteCount {
5660 return switch (try frameType(src)) {
5761 .zstandard => decodeZStandardFrame(dest, src, verify_checksum),
5862 .skippable => ReadWriteCount{
......@@ -100,7 +104,7 @@ pub const DecodeState = struct {
100104 src: []const u8,
101105 literals: LiteralsSection,
102106 sequences_header: SequencesSection.Header,
103 ) !usize {
107 ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize {
104108 if (literals.huffman_tree) |tree| {
105109 self.huffman_tree = tree;
106110 } else if (literals.header.block_type == .treeless and self.huffman_tree == null) {
......@@ -145,7 +149,7 @@ pub const DecodeState = struct {
145149
146150 /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream`
147151 /// if `bit_reader` does not contain enough bits.
148 pub fn readInitialFseState(self: *DecodeState, bit_reader: anytype) !void {
152 pub fn readInitialFseState(self: *DecodeState, bit_reader: *ReverseBitReader) error{EndOfStream}!void {
149153 self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log);
150154 self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log);
151155 self.match.state = try bit_reader.readBitsNoEof(u9, self.match.accuracy_log);
......@@ -169,7 +173,11 @@ pub const DecodeState = struct {
169173
170174 const DataType = enum { offset, match, literal };
171175
172 fn updateState(self: *DecodeState, comptime choice: DataType, bit_reader: anytype) !void {
176 fn updateState(
177 self: *DecodeState,
178 comptime choice: DataType,
179 bit_reader: *ReverseBitReader,
180 ) error{ MalformedFseBits, EndOfStream }!void {
173181 switch (@field(self, @tagName(choice)).table) {
174182 .rle => {},
175183 .fse => |table| {
......@@ -185,17 +193,27 @@ pub const DecodeState = struct {
185193 }
186194 }
187195
196 const FseTableError = error{
197 MalformedFseTable,
198 MalformedAccuracyLog,
199 RepeatModeFirst,
200 EndOfStream,
201 };
202
188203 fn updateFseTable(
189204 self: *DecodeState,
190205 src: []const u8,
191206 comptime choice: DataType,
192207 mode: SequencesSection.Header.Mode,
193 ) !usize {
208 ) FseTableError!usize {
194209 const field_name = @tagName(choice);
195210 switch (mode) {
196211 .predefined => {
197 @field(self, field_name).accuracy_log = @field(types.compressed_block.default_accuracy_log, field_name);
198 @field(self, field_name).table = @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table");
212 @field(self, field_name).accuracy_log =
213 @field(types.compressed_block.default_accuracy_log, field_name);
214
215 @field(self, field_name).table =
216 @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table");
199217 return 0;
200218 },
201219 .rle => {
......@@ -214,9 +232,11 @@ pub const DecodeState = struct {
214232 @field(types.compressed_block.table_accuracy_log_max, field_name),
215233 @field(self, field_name ++ "_fse_buffer"),
216234 );
217 @field(self, field_name).table = .{ .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size] };
235 @field(self, field_name).table = .{
236 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],
237 };
218238 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);
219 return std.math.cast(usize, counting_reader.bytes_read) orelse return error.MalformedFseTable;
239 return std.math.cast(usize, counting_reader.bytes_read) orelse error.MalformedFseTable;
220240 },
221241 .repeat => return if (self.fse_tables_undefined) error.RepeatModeFirst else 0,
222242 }
......@@ -228,7 +248,10 @@ pub const DecodeState = struct {
228248 offset: u32,
229249 };
230250
231 fn nextSequence(self: *DecodeState, bit_reader: anytype) !Sequence {
251 fn nextSequence(
252 self: *DecodeState,
253 bit_reader: *ReverseBitReader,
254 ) error{ OffsetCodeTooLarge, EndOfStream }!Sequence {
232255 const raw_code = self.getCode(.offset);
233256 const offset_code = std.math.cast(u5, raw_code) orelse {
234257 return error.OffsetCodeTooLarge;
......@@ -272,7 +295,7 @@ pub const DecodeState = struct {
272295 write_pos: usize,
273296 literals: LiteralsSection,
274297 sequence: Sequence,
275 ) !void {
298 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
276299 if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence;
277300
278301 try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length);
......@@ -288,16 +311,23 @@ pub const DecodeState = struct {
288311 dest: *RingBuffer,
289312 literals: LiteralsSection,
290313 sequence: Sequence,
291 ) !void {
314 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
292315 if (sequence.offset > dest.data.len) return error.MalformedSequence;
293316
294317 try self.decodeLiteralsRingBuffer(dest, literals, sequence.literal_length);
295 const copy_slice = dest.sliceAt(dest.write_index + dest.data.len - sequence.offset, sequence.match_length);
318 const copy_start = dest.write_index + dest.data.len - sequence.offset;
319 const copy_slice = dest.sliceAt(copy_start, sequence.match_length);
296320 // TODO: would std.mem.copy and figuring out dest slice be better/faster?
297321 for (copy_slice.first) |b| dest.writeAssumeCapacity(b);
298322 for (copy_slice.second) |b| dest.writeAssumeCapacity(b);
299323 }
300324
325 const DecodeSequenceError = error{
326 OffsetCodeTooLarge,
327 EndOfStream,
328 MalformedSequence,
329 MalformedFseBits,
330 } || DecodeLiteralsError;
301331 /// Decode one sequence from `bit_reader` into `dest`, written starting at
302332 /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns
303333 /// `error.MalformedSequence` error if the decompressed sequence would be longer
......@@ -311,10 +341,10 @@ pub const DecodeState = struct {
311341 dest: []u8,
312342 write_pos: usize,
313343 literals: LiteralsSection,
314 bit_reader: anytype,
344 bit_reader: *ReverseBitReader,
315345 sequence_size_limit: usize,
316346 last_sequence: bool,
317 ) !usize {
347 ) DecodeSequenceError!usize {
318348 const sequence = try self.nextSequence(bit_reader);
319349 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
320350 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
......@@ -336,7 +366,7 @@ pub const DecodeState = struct {
336366 bit_reader: anytype,
337367 sequence_size_limit: usize,
338368 last_sequence: bool,
339 ) !usize {
369 ) DecodeSequenceError!usize {
340370 const sequence = try self.nextSequence(bit_reader);
341371 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
342372 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
......@@ -350,26 +380,63 @@ pub const DecodeState = struct {
350380 return sequence_length;
351381 }
352382
353 fn nextLiteralMultiStream(self: *DecodeState, literals: LiteralsSection) !void {
383 fn nextLiteralMultiStream(
384 self: *DecodeState,
385 literals: LiteralsSection,
386 ) error{BitStreamHasNoStartBit}!void {
354387 self.literal_stream_index += 1;
355388 try self.initLiteralStream(literals.streams.four[self.literal_stream_index]);
356389 }
357390
358 fn initLiteralStream(self: *DecodeState, bytes: []const u8) !void {
391 fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
359392 try self.literal_stream_reader.init(bytes);
360393 }
361394
395 const LiteralBitsError = error{
396 BitStreamHasNoStartBit,
397 UnexpectedEndOfLiteralStream,
398 };
399 fn readLiteralsBits(
400 self: *DecodeState,
401 comptime T: type,
402 bit_count_to_read: usize,
403 literals: LiteralsSection,
404 ) LiteralBitsError!T {
405 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {
406 if (literals.streams == .four and self.literal_stream_index < 3) {
407 try self.nextLiteralMultiStream(literals);
408 break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch
409 return error.UnexpectedEndOfLiteralStream;
410 } else {
411 return error.UnexpectedEndOfLiteralStream;
412 }
413 };
414 }
415
416 const DecodeLiteralsError = error{
417 MalformedLiteralsLength,
418 PrefixNotFound,
419 } || LiteralBitsError;
420
362421 /// Decode `len` bytes of literals into `dest`. `literals` should be the
363422 /// `LiteralsSection` that was passed to `prepare()`. Returns
364423 /// `error.MalformedLiteralsLength` if the number of literal bytes decoded by
365424 /// `self` plus `len` is greater than the regenerated size of `literals`.
366425 /// Returns `error.UnexpectedEndOfLiteralStream` and `error.PrefixNotFound` if
367426 /// there are problems decoding Huffman compressed literals.
368 pub fn decodeLiteralsSlice(self: *DecodeState, dest: []u8, literals: LiteralsSection, len: usize) !void {
369 if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength;
427 pub fn decodeLiteralsSlice(
428 self: *DecodeState,
429 dest: []u8,
430 literals: LiteralsSection,
431 len: usize,
432 ) DecodeLiteralsError!void {
433 if (self.literal_written_count + len > literals.header.regenerated_size)
434 return error.MalformedLiteralsLength;
435
370436 switch (literals.header.block_type) {
371437 .raw => {
372 const literal_data = literals.streams.one[self.literal_written_count .. self.literal_written_count + len];
438 const literals_end = self.literal_written_count + len;
439 const literal_data = literals.streams.one[self.literal_written_count..literals_end];
373440 std.mem.copy(u8, dest, literal_data);
374441 self.literal_written_count += len;
375442 },
......@@ -395,15 +462,7 @@ pub const DecodeState = struct {
395462 while (i < len) : (i += 1) {
396463 var prefix: u16 = 0;
397464 while (true) {
398 const new_bits = self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |err|
399 switch (err) {
400 error.EndOfStream => if (literals.streams == .four and self.literal_stream_index < 3) bits: {
401 try self.nextLiteralMultiStream(literals);
402 break :bits try self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read);
403 } else {
404 return error.UnexpectedEndOfLiteralStream;
405 },
406 };
465 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals);
407466 prefix <<= bit_count_to_read;
408467 prefix |= new_bits;
409468 bits_read += bit_count_to_read;
......@@ -434,11 +493,19 @@ pub const DecodeState = struct {
434493 }
435494
436495 /// Decode literals into `dest`; see `decodeLiteralsSlice()`.
437 pub fn decodeLiteralsRingBuffer(self: *DecodeState, dest: *RingBuffer, literals: LiteralsSection, len: usize) !void {
438 if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength;
496 pub fn decodeLiteralsRingBuffer(
497 self: *DecodeState,
498 dest: *RingBuffer,
499 literals: LiteralsSection,
500 len: usize,
501 ) DecodeLiteralsError!void {
502 if (self.literal_written_count + len > literals.header.regenerated_size)
503 return error.MalformedLiteralsLength;
504
439505 switch (literals.header.block_type) {
440506 .raw => {
441 const literal_data = literals.streams.one[self.literal_written_count .. self.literal_written_count + len];
507 const literals_end = self.literal_written_count + len;
508 const literal_data = literals.streams.one[self.literal_written_count..literals_end];
442509 dest.writeSliceAssumeCapacity(literal_data);
443510 self.literal_written_count += len;
444511 },
......@@ -464,15 +531,7 @@ pub const DecodeState = struct {
464531 while (i < len) : (i += 1) {
465532 var prefix: u16 = 0;
466533 while (true) {
467 const new_bits = self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |err|
468 switch (err) {
469 error.EndOfStream => if (literals.streams == .four and self.literal_stream_index < 3) bits: {
470 try self.nextLiteralMultiStream(literals);
471 break :bits try self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read);
472 } else {
473 return error.UnexpectedEndOfLiteralStream;
474 },
475 };
534 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals);
476535 prefix <<= bit_count_to_read;
477536 prefix |= new_bits;
478537 bits_read += bit_count_to_read;
......@@ -514,6 +573,11 @@ const literal_table_size_max = 1 << types.compressed_block.table_accuracy_log_ma
514573const match_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match;
515574const offset_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match;
516575
576const FrameError = error{
577 DictionaryIdFlagUnsupported,
578 ChecksumFailure,
579} || InvalidBit || DecodeBlockError;
580
517581/// Decode a Zstandard frame from `src` into `dest`, returning the number of
518582/// bytes read from `src` and written to `dest`; if the frame does not declare
519583/// its decompressed content size `error.UnknownContentSizeUnsupported` is
......@@ -521,7 +585,11 @@ const offset_table_size_max = 1 << types.compressed_block.table_accuracy_log_max
521585/// dictionary, and `error.ChecksumFailure` if `verify_checksum` is `true` and
522586/// the frame contains a checksum that does not match the checksum computed from
523587/// the decompressed frame.
524pub fn decodeZStandardFrame(dest: []u8, src: []const u8, verify_checksum: bool) !ReadWriteCount {
588pub fn decodeZStandardFrame(
589 dest: []u8,
590 src: []const u8,
591 verify_checksum: bool,
592) (error{ UnknownContentSizeUnsupported, ContentTooLarge } || FrameError)!ReadWriteCount {
525593 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
526594 var consumed_count: usize = 4;
527595
......@@ -530,13 +598,11 @@ pub fn decodeZStandardFrame(dest: []u8, src: []const u8, verify_checksum: bool)
530598 if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported;
531599
532600 const content_size = frame_header.content_size orelse return error.UnknownContentSizeUnsupported;
533 // const window_size = frameWindowSize(header) orelse return error.WindowSizeUnknown;
534601 if (dest.len < content_size) return error.ContentTooLarge;
535602
536603 const should_compute_checksum = frame_header.descriptor.content_checksum_flag and verify_checksum;
537604 var hash_state = if (should_compute_checksum) std.hash.XxHash64.init(0) else undefined;
538605
539 // TODO: block_maximum_size should be @min(1 << 17, window_size);
540606 const written_count = try decodeFrameBlocks(
541607 dest,
542608 src[consumed_count..],
......@@ -567,7 +633,7 @@ pub fn decodeZStandardFrameAlloc(
567633 src: []const u8,
568634 verify_checksum: bool,
569635 window_size_max: usize,
570) ![]u8 {
636) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory } || FrameError)![]u8 {
571637 var result = std.ArrayList(u8).init(allocator);
572638 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
573639 var consumed_count: usize = 4;
......@@ -628,7 +694,7 @@ pub fn decodeZStandardFrameAlloc(
628694 block_header = decodeBlockHeader(src[consumed_count..][0..3]);
629695 consumed_count += 3;
630696 }) {
631 if (block_header.block_size > block_size_maximum) return error.CompressedBlockSizeOverMaximum;
697 if (block_header.block_size > block_size_maximum) return error.BlockSizeOverMaximum;
632698 const written_size = try decodeBlockRingBuffer(
633699 &ring_buffer,
634700 src[consumed_count..],
......@@ -637,7 +703,7 @@ pub fn decodeZStandardFrameAlloc(
637703 &consumed_count,
638704 block_size_maximum,
639705 );
640 if (written_size > block_size_maximum) return error.DecompressedBlockSizeOverMaximum;
706 if (written_size > block_size_maximum) return error.BlockSizeOverMaximum;
641707 const written_slice = ring_buffer.sliceLast(written_size);
642708 try result.appendSlice(written_slice.first);
643709 try result.appendSlice(written_slice.second);
......@@ -650,8 +716,21 @@ pub fn decodeZStandardFrameAlloc(
650716 return result.toOwnedSlice();
651717}
652718
719const DecodeBlockError = error{
720 BlockSizeOverMaximum,
721 MalformedBlockSize,
722 ReservedBlock,
723 MalformedRleBlock,
724 MalformedCompressedBlock,
725};
726
653727/// Convenience wrapper for decoding all blocks in a frame; see `decodeBlock()`.
654pub fn decodeFrameBlocks(dest: []u8, src: []const u8, consumed_count: *usize, hash: ?*std.hash.XxHash64) !usize {
728pub fn decodeFrameBlocks(
729 dest: []u8,
730 src: []const u8,
731 consumed_count: *usize,
732 hash: ?*std.hash.XxHash64,
733) DecodeBlockError!usize {
655734 // These tables take 7680 bytes
656735 var literal_fse_data: [literal_table_size_max]Table.Fse = undefined;
657736 var match_fse_data: [match_table_size_max]Table.Fse = undefined;
......@@ -702,7 +781,12 @@ pub fn decodeFrameBlocks(dest: []u8, src: []const u8, consumed_count: *usize, ha
702781 return written_count;
703782}
704783
705fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
784fn decodeRawBlock(
785 dest: []u8,
786 src: []const u8,
787 block_size: u21,
788 consumed_count: *usize,
789) error{MalformedBlockSize}!usize {
706790 if (src.len < block_size) return error.MalformedBlockSize;
707791 const data = src[0..block_size];
708792 std.mem.copy(u8, dest, data);
......@@ -710,7 +794,12 @@ fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:
710794 return block_size;
711795}
712796
713fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
797fn decodeRawBlockRingBuffer(
798 dest: *RingBuffer,
799 src: []const u8,
800 block_size: u21,
801 consumed_count: *usize,
802) error{MalformedBlockSize}!usize {
714803 if (src.len < block_size) return error.MalformedBlockSize;
715804 const data = src[0..block_size];
716805 dest.writeSliceAssumeCapacity(data);
......@@ -718,7 +807,12 @@ fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21,
718807 return block_size;
719808}
720809
721fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
810fn decodeRleBlock(
811 dest: []u8,
812 src: []const u8,
813 block_size: u21,
814 consumed_count: *usize,
815) error{MalformedRleBlock}!usize {
722816 if (src.len < 1) return error.MalformedRleBlock;
723817 var write_pos: usize = 0;
724818 while (write_pos < block_size) : (write_pos += 1) {
......@@ -728,7 +822,12 @@ fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:
728822 return block_size;
729823}
730824
731fn decodeRleBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) !usize {
825fn decodeRleBlockRingBuffer(
826 dest: *RingBuffer,
827 src: []const u8,
828 block_size: u21,
829 consumed_count: *usize,
830) error{MalformedRleBlock}!usize {
732831 if (src.len < 1) return error.MalformedRleBlock;
733832 var write_pos: usize = 0;
734833 while (write_pos < block_size) : (write_pos += 1) {
......@@ -749,7 +848,7 @@ pub fn decodeBlock(
749848 decode_state: *DecodeState,
750849 consumed_count: *usize,
751850 written_count: usize,
752) !usize {
851) DecodeBlockError!usize {
753852 const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB
754853 const block_size = block_header.block_size;
755854 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
......@@ -759,31 +858,33 @@ pub fn decodeBlock(
759858 .compressed => {
760859 if (src.len < block_size) return error.MalformedBlockSize;
761860 var bytes_read: usize = 0;
762 const literals = try decodeLiteralsSection(src, &bytes_read);
763 const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read);
861 const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock;
862 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
863 return error.MalformedCompressedBlock;
764864
765 bytes_read += try decode_state.prepare(src[bytes_read..], literals, sequences_header);
865 bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch
866 return error.MalformedCompressedBlock;
766867
767868 var bytes_written: usize = 0;
768869 if (sequences_header.sequence_count > 0) {
769870 const bit_stream_bytes = src[bytes_read..block_size];
770871 var bit_stream: ReverseBitReader = undefined;
771 try bit_stream.init(bit_stream_bytes);
872 bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock;
772873
773 try decode_state.readInitialFseState(&bit_stream);
874 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
774875
775876 var sequence_size_limit = block_size_max;
776877 var i: usize = 0;
777878 while (i < sequences_header.sequence_count) : (i += 1) {
778879 const write_pos = written_count + bytes_written;
779 const decompressed_size = try decode_state.decodeSequenceSlice(
880 const decompressed_size = decode_state.decodeSequenceSlice(
780881 dest,
781882 write_pos,
782883 literals,
783884 &bit_stream,
784885 sequence_size_limit,
785886 i == sequences_header.sequence_count - 1,
786 );
887 ) catch return error.MalformedCompressedBlock;
787888 bytes_written += decompressed_size;
788889 sequence_size_limit -= decompressed_size;
789890 }
......@@ -793,7 +894,8 @@ pub fn decodeBlock(
793894
794895 if (decode_state.literal_written_count < literals.header.regenerated_size) {
795896 const len = literals.header.regenerated_size - decode_state.literal_written_count;
796 try decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], literals, len);
897 decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], literals, len) catch
898 return error.MalformedCompressedBlock;
797899 bytes_written += len;
798900 }
799901
......@@ -802,7 +904,7 @@ pub fn decodeBlock(
802904 consumed_count.* += bytes_read;
803905 return bytes_written;
804906 },
805 .reserved => return error.FrameContainsReservedBlock,
907 .reserved => return error.ReservedBlock,
806908 }
807909}
808910
......@@ -816,7 +918,7 @@ pub fn decodeBlockRingBuffer(
816918 decode_state: *DecodeState,
817919 consumed_count: *usize,
818920 block_size_max: usize,
819) !usize {
921) DecodeBlockError!usize {
820922 const block_size = block_header.block_size;
821923 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
822924 switch (block_header.block_type) {
......@@ -825,29 +927,31 @@ pub fn decodeBlockRingBuffer(
825927 .compressed => {
826928 if (src.len < block_size) return error.MalformedBlockSize;
827929 var bytes_read: usize = 0;
828 const literals = try decodeLiteralsSection(src, &bytes_read);
829 const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read);
930 const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock;
931 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
932 return error.MalformedCompressedBlock;
830933
831 bytes_read += try decode_state.prepare(src[bytes_read..], literals, sequences_header);
934 bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch
935 return error.MalformedCompressedBlock;
832936
833937 var bytes_written: usize = 0;
834938 if (sequences_header.sequence_count > 0) {
835939 const bit_stream_bytes = src[bytes_read..block_size];
836940 var bit_stream: ReverseBitReader = undefined;
837 try bit_stream.init(bit_stream_bytes);
941 bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock;
838942
839 try decode_state.readInitialFseState(&bit_stream);
943 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
840944
841945 var sequence_size_limit = block_size_max;
842946 var i: usize = 0;
843947 while (i < sequences_header.sequence_count) : (i += 1) {
844 const decompressed_size = try decode_state.decodeSequenceRingBuffer(
948 const decompressed_size = decode_state.decodeSequenceRingBuffer(
845949 dest,
846950 literals,
847951 &bit_stream,
848952 sequence_size_limit,
849953 i == sequences_header.sequence_count - 1,
850 );
954 ) catch return error.MalformedCompressedBlock;
851955 bytes_written += decompressed_size;
852956 sequence_size_limit -= decompressed_size;
853957 }
......@@ -857,7 +961,8 @@ pub fn decodeBlockRingBuffer(
857961
858962 if (decode_state.literal_written_count < literals.header.regenerated_size) {
859963 const len = literals.header.regenerated_size - decode_state.literal_written_count;
860 try decode_state.decodeLiteralsRingBuffer(dest, literals, len);
964 decode_state.decodeLiteralsRingBuffer(dest, literals, len) catch
965 return error.MalformedCompressedBlock;
861966 bytes_written += len;
862967 }
863968
......@@ -866,7 +971,7 @@ pub fn decodeBlockRingBuffer(
866971 consumed_count.* += bytes_read;
867972 return bytes_written;
868973 },
869 .reserved => return error.FrameContainsReservedBlock,
974 .reserved => return error.ReservedBlock,
870975 }
871976}
872977
......@@ -901,9 +1006,10 @@ pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 {
9011006 } else return header.content_size;
9021007}
9031008
1009const InvalidBit = error{ UnusedBitSet, ReservedBitSet };
9041010/// Decode the header of a Zstandard frame. Returns `error.UnusedBitSet` or
9051011/// `error.ReservedBitSet` if the corresponding bits are sets.
906pub fn decodeZStandardHeader(src: []const u8, consumed_count: ?*usize) !frame.ZStandard.Header {
1012pub fn decodeZStandardHeader(src: []const u8, consumed_count: ?*usize) InvalidBit!frame.ZStandard.Header {
9071013 const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, src[0]);
9081014
9091015 if (descriptor.unused) return error.UnusedBitSet;
......@@ -958,7 +1064,10 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {
9581064
9591065/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
9601066/// number of bytes the section uses.
961pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsSection {
1067pub fn decodeLiteralsSection(
1068 src: []const u8,
1069 consumed_count: *usize,
1070) (error{ MalformedLiteralsHeader, MalformedLiteralsSection } || DecodeHuffmanError)!LiteralsSection {
9621071 var bytes_read: usize = 0;
9631072 const header = try decodeLiteralsHeader(src, &bytes_read);
9641073 switch (header.block_type) {
......@@ -1032,7 +1141,13 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS
10321141 }
10331142}
10341143
1035fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.HuffmanTree {
1144const DecodeHuffmanError = error{
1145 MalformedHuffmanTree,
1146 MalformedFseTable,
1147 MalformedAccuracyLog,
1148};
1149
1150fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError!LiteralsSection.HuffmanTree {
10361151 var bytes_read: usize = 0;
10371152 bytes_read += 1;
10381153 if (src.len == 0) return error.MalformedHuffmanTree;
......@@ -1049,22 +1164,25 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H
10491164 var bit_reader = bitReader(counting_reader.reader());
10501165
10511166 var entries: [1 << 6]Table.Fse = undefined;
1052 const table_size = try decodeFseTable(&bit_reader, 256, 6, &entries);
1167 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
1168 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
1169 error.EndOfStream => return error.MalformedFseTable,
1170 };
10531171 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
10541172
10551173 const start_index = std.math.cast(usize, 1 + counting_reader.bytes_read) orelse return error.MalformedHuffmanTree;
10561174 var huff_data = src[start_index .. compressed_size + 1];
10571175 var huff_bits: ReverseBitReader = undefined;
1058 try huff_bits.init(huff_data);
1176 huff_bits.init(huff_data) catch return error.MalformedHuffmanTree;
10591177
10601178 var i: usize = 0;
1061 var even_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log);
1062 var odd_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log);
1179 var even_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
1180 var odd_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
10631181
10641182 while (i < 255) {
10651183 const even_data = entries[even_state];
10661184 var read_bits: usize = 0;
1067 const even_bits = try huff_bits.readBits(u32, even_data.bits, &read_bits);
1185 const even_bits = huff_bits.readBits(u32, even_data.bits, &read_bits) catch unreachable;
10681186 weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree;
10691187 i += 1;
10701188 if (read_bits < even_data.bits) {
......@@ -1076,7 +1194,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H
10761194
10771195 read_bits = 0;
10781196 const odd_data = entries[odd_state];
1079 const odd_bits = try huff_bits.readBits(u32, odd_data.bits, &read_bits);
1197 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;
10801198 weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree;
10811199 i += 1;
10821200 if (read_bits < odd_data.bits) {
......@@ -1177,8 +1295,8 @@ fn lessThanByWeight(
11771295}
11781296
11791297/// Decode a literals section header.
1180pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSection.Header {
1181 if (src.len == 0) return error.MalformedLiteralsSection;
1298pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{MalformedLiteralsHeader}!LiteralsSection.Header {
1299 if (src.len == 0) return error.MalformedLiteralsHeader;
11821300 const byte0 = src[0];
11831301 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);
11841302 const size_format = @intCast(u2, (byte0 & 0b1100) >> 2);
......@@ -1243,8 +1361,11 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSe
12431361}
12441362
12451363/// Decode a sequences section header.
1246pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !SequencesSection.Header {
1247 if (src.len == 0) return error.MalformedSequencesSection;
1364pub fn decodeSequencesHeader(
1365 src: []const u8,
1366 consumed_count: *usize,
1367) error{ MalformedSequencesHeader, ReservedBitSet }!SequencesSection.Header {
1368 if (src.len == 0) return error.MalformedSequencesHeader;
12481369 var sequence_count: u24 = undefined;
12491370
12501371 var bytes_read: usize = 0;
......@@ -1262,16 +1383,16 @@ pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !Sequences
12621383 sequence_count = byte0;
12631384 bytes_read += 1;
12641385 } else if (byte0 < 255) {
1265 if (src.len < 2) return error.MalformedSequencesSection;
1386 if (src.len < 2) return error.MalformedSequencesHeader;
12661387 sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1];
12671388 bytes_read += 2;
12681389 } else {
1269 if (src.len < 3) return error.MalformedSequencesSection;
1390 if (src.len < 3) return error.MalformedSequencesHeader;
12701391 sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00;
12711392 bytes_read += 3;
12721393 }
12731394
1274 if (src.len < bytes_read + 1) return error.MalformedSequencesSection;
1395 if (src.len < bytes_read + 1) return error.MalformedSequencesHeader;
12751396 const compression_modes = src[bytes_read];
12761397 bytes_read += 1;
12771398
......@@ -1441,17 +1562,17 @@ pub const ReverseBitReader = struct {
14411562 byte_reader: ReversedByteReader,
14421563 bit_reader: std.io.BitReader(.Big, ReversedByteReader.Reader),
14431564
1444 pub fn init(self: *ReverseBitReader, bytes: []const u8) !void {
1565 pub fn init(self: *ReverseBitReader, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
14451566 self.byte_reader = ReversedByteReader.init(bytes);
14461567 self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader());
14471568 while (0 == self.readBitsNoEof(u1, 1) catch return error.BitStreamHasNoStartBit) {}
14481569 }
14491570
1450 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) !U {
1571 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) error{EndOfStream}!U {
14511572 return self.bit_reader.readBitsNoEof(U, num_bits);
14521573 }
14531574
1454 pub fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) !U {
1575 pub fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) error{}!U {
14551576 return try self.bit_reader.readBits(U, num_bits, out_bits);
14561577 }
14571578
lib/std/compress/zstandard/types.zig+1-1
......@@ -92,7 +92,7 @@ pub const compressed_block = struct {
9292 index: usize,
9393 };
9494
95 pub fn query(self: HuffmanTree, index: usize, prefix: u16) !Result {
95 pub fn query(self: HuffmanTree, index: usize, prefix: u16) error{PrefixNotFound}!Result {
9696 var node = self.nodes[index];
9797 const weight = node.weight;
9898 var i: usize = index;