authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-02 16:19:13+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log5723291444116419440a187adcfa5ecb9557544e
treee734bf6e22afa99182422fae8e01895cef65604d
parent947ad3e26816270fa11298dd5f118e50126a1320

std.compress.zstandard: add `decodeBlockReader`


1 files changed, 463 insertions(+), 361 deletions(-)

lib/std/compress/zstandard/decompress.zig+463-361
...@@ -14,29 +14,18 @@ fn readVarInt(comptime T: type, bytes: []const u8) T {...@@ -14,29 +14,18 @@ fn readVarInt(comptime T: type, bytes: []const u8) T {
14 return std.mem.readVarInt(T, bytes, .Little);14 return std.mem.readVarInt(T, bytes, .Little);
15}15}
1616
17fn isSkippableMagic(magic: u32) bool {17pub fn isSkippableMagic(magic: u32) bool {
18 return frame.Skippable.magic_number_min <= magic and magic <= frame.Skippable.magic_number_max;18 return frame.Skippable.magic_number_min <= magic and magic <= frame.Skippable.magic_number_max;
19}19}
2020
21/// Returns the decompressed size of the frame at the start of `src`. Returns 021/// Returns the kind of frame at the beginning of `src`.
22/// if the the frame is skippable, `null` for Zstanndard frames that do not22///
23/// declare their content size. Returns `UnusedBitSet` and `ReservedBitSet`23/// Errors:
24/// errors if the respective bits of the the frame descriptor are set.24/// - returns `error.BadMagic` if `source` begins with bytes not equal to the
25pub fn getFrameDecompressedSize(src: []const u8) (InvalidBit || error{BadMagic})!?u64 {25/// Zstandard frame magic number, or outside the range of magic numbers for
26 switch (try frameType(src)) {26/// skippable frames.
27 .zstandard => {27pub fn decodeFrameType(source: anytype) !frame.Kind {
28 const header = try decodeZStandardHeader(src[4..], null);28 const magic = try source.readIntLittle(u32);
29 return header.content_size;
30 },
31 .skippable => return 0,
32 }
33}
34
35/// Returns the kind of frame at the beginning of `src`. Returns `BadMagic` if
36/// `src` begin with bytes not equal to the Zstandard frame magic number, or
37/// outside the range of magic numbers for skippable frames.
38pub fn frameType(src: []const u8) error{BadMagic}!frame.Kind {
39 const magic = readInt(u32, src[0..4]);
40 return if (magic == frame.ZStandard.magic_number)29 return if (magic == frame.ZStandard.magic_number)
41 .zstandard30 .zstandard
42 else if (isSkippableMagic(magic))31 else if (isSkippableMagic(magic))
...@@ -52,15 +41,21 @@ const ReadWriteCount = struct {...@@ -52,15 +41,21 @@ const ReadWriteCount = struct {
5241
53/// Decodes the frame at the start of `src` into `dest`. Returns the number of42/// Decodes the frame at the start of `src` into `dest`. Returns the number of
54/// bytes read from `src` and written to `dest`.43/// bytes read from `src` and written to `dest`.
44///
45/// Errors:
46/// - returns `error.UnknownContentSizeUnsupported`
47/// - returns `error.ContentTooLarge`
48/// - returns `error.BadMagic`
55pub fn decodeFrame(49pub fn decodeFrame(
56 dest: []u8,50 dest: []u8,
57 src: []const u8,51 src: []const u8,
58 verify_checksum: bool,52 verify_checksum: bool,
59) (error{ UnknownContentSizeUnsupported, ContentTooLarge, BadMagic } || FrameError)!ReadWriteCount {53) !ReadWriteCount {
60 return switch (try frameType(src)) {54 var fbs = std.io.fixedBufferStream(src);
55 return switch (try decodeFrameType(fbs.reader())) {
61 .zstandard => decodeZStandardFrame(dest, src, verify_checksum),56 .zstandard => decodeZStandardFrame(dest, src, verify_checksum),
62 .skippable => ReadWriteCount{57 .skippable => ReadWriteCount{
63 .read_count = skippableFrameSize(src[0..8]) + 8,58 .read_count = try fbs.reader().readIntLittle(u32) + 8,
64 .write_count = 0,59 .write_count = 0,
65 },60 },
66 };61 };
...@@ -97,16 +92,52 @@ pub const DecodeState = struct {...@@ -97,16 +92,52 @@ pub const DecodeState = struct {
97 };92 };
98 }93 }
9994
95 pub fn init(
96 literal_fse_buffer: []Table.Fse,
97 match_fse_buffer: []Table.Fse,
98 offset_fse_buffer: []Table.Fse,
99 ) DecodeState {
100 return DecodeState{
101 .repeat_offsets = .{
102 types.compressed_block.start_repeated_offset_1,
103 types.compressed_block.start_repeated_offset_2,
104 types.compressed_block.start_repeated_offset_3,
105 },
106
107 .offset = undefined,
108 .match = undefined,
109 .literal = undefined,
110
111 .literal_fse_buffer = literal_fse_buffer,
112 .match_fse_buffer = match_fse_buffer,
113 .offset_fse_buffer = offset_fse_buffer,
114
115 .fse_tables_undefined = true,
116
117 .literal_written_count = 0,
118 .literal_header = undefined,
119 .literal_streams = undefined,
120 .literal_stream_reader = undefined,
121 .literal_stream_index = undefined,
122 .huffman_tree = null,
123 };
124 }
125
100 /// Prepare the decoder to decode a compressed block. Loads the literals126 /// Prepare the decoder to decode a compressed block. Loads the literals
101 /// stream and Huffman tree from `literals` and reads the FSE tables from `src`.127 /// stream and Huffman tree from `literals` and reads the FSE tables from
102 /// Returns `error.BitStreamHasNoStartBit` if the (reversed) literal bitstream's128 /// `source`.
103 /// first byte does not have any bits set.129 ///
130 /// Errors:
131 /// - returns `error.BitStreamHasNoStartBit` if the (reversed) literal bitstream's
132 /// first byte does not have any bits set.
133 /// - returns `error.TreelessLiteralsFirst` `literals` is a treeless literals section
134 /// and the decode state does not have a Huffman tree from a previous block.
104 pub fn prepare(135 pub fn prepare(
105 self: *DecodeState,136 self: *DecodeState,
106 src: []const u8,137 source: anytype,
107 literals: LiteralsSection,138 literals: LiteralsSection,
108 sequences_header: SequencesSection.Header,139 sequences_header: SequencesSection.Header,
109 ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize {140 ) !void {
110 self.literal_written_count = 0;141 self.literal_written_count = 0;
111 self.literal_header = literals.header;142 self.literal_header = literals.header;
112 self.literal_streams = literals.streams;143 self.literal_streams = literals.streams;
...@@ -129,28 +160,11 @@ pub const DecodeState = struct {...@@ -129,28 +160,11 @@ pub const DecodeState = struct {
129 }160 }
130161
131 if (sequences_header.sequence_count > 0) {162 if (sequences_header.sequence_count > 0) {
132 var bytes_read = try self.updateFseTable(163 try self.updateFseTable(source, .literal, sequences_header.literal_lengths);
133 src,164 try self.updateFseTable(source, .offset, sequences_header.offsets);
134 .literal,165 try self.updateFseTable(source, .match, sequences_header.match_lengths);
135 sequences_header.literal_lengths,
136 );
137
138 bytes_read += try self.updateFseTable(
139 src[bytes_read..],
140 .offset,
141 sequences_header.offsets,
142 );
143
144 bytes_read += try self.updateFseTable(
145 src[bytes_read..],
146 .match,
147 sequences_header.match_lengths,
148 );
149 self.fse_tables_undefined = false;166 self.fse_tables_undefined = false;
150
151 return bytes_read;
152 }167 }
153 return 0;
154 }168 }
155169
156 /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream`170 /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream`
...@@ -208,10 +222,10 @@ pub const DecodeState = struct {...@@ -208,10 +222,10 @@ pub const DecodeState = struct {
208222
209 fn updateFseTable(223 fn updateFseTable(
210 self: *DecodeState,224 self: *DecodeState,
211 src: []const u8,225 source: anytype,
212 comptime choice: DataType,226 comptime choice: DataType,
213 mode: SequencesSection.Header.Mode,227 mode: SequencesSection.Header.Mode,
214 ) FseTableError!usize {228 ) !void {
215 const field_name = @tagName(choice);229 const field_name = @tagName(choice);
216 switch (mode) {230 switch (mode) {
217 .predefined => {231 .predefined => {
...@@ -220,17 +234,13 @@ pub const DecodeState = struct {...@@ -220,17 +234,13 @@ pub const DecodeState = struct {
220234
221 @field(self, field_name).table =235 @field(self, field_name).table =
222 @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table");236 @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table");
223 return 0;
224 },237 },
225 .rle => {238 .rle => {
226 @field(self, field_name).accuracy_log = 0;239 @field(self, field_name).accuracy_log = 0;
227 @field(self, field_name).table = .{ .rle = src[0] };240 @field(self, field_name).table = .{ .rle = try source.readByte() };
228 return 1;
229 },241 },
230 .fse => {242 .fse => {
231 var stream = std.io.fixedBufferStream(src);243 var bit_reader = bitReader(source);
232 var counting_reader = std.io.countingReader(stream.reader());
233 var bit_reader = bitReader(counting_reader.reader());
234244
235 const table_size = try decodeFseTable(245 const table_size = try decodeFseTable(
236 &bit_reader,246 &bit_reader,
...@@ -242,9 +252,8 @@ pub const DecodeState = struct {...@@ -242,9 +252,8 @@ pub const DecodeState = struct {
242 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],252 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],
243 };253 };
244 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);254 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);
245 return std.math.cast(usize, counting_reader.bytes_read) orelse error.MalformedFseTable;
246 },255 },
247 .repeat => return if (self.fse_tables_undefined) error.RepeatModeFirst else 0,256 .repeat => if (self.fse_tables_undefined) return error.RepeatModeFirst,
248 }257 }
249 }258 }
250259
...@@ -462,11 +471,15 @@ pub const DecodeState = struct {...@@ -462,11 +471,15 @@ pub const DecodeState = struct {
462 while (i < len) : (i += 1) {471 while (i < len) : (i += 1) {
463 var prefix: u16 = 0;472 var prefix: u16 = 0;
464 while (true) {473 while (true) {
465 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read);474 const new_bits = self.readLiteralsBits(u16, bit_count_to_read) catch |err| {
475 return err;
476 };
466 prefix <<= bit_count_to_read;477 prefix <<= bit_count_to_read;
467 prefix |= new_bits;478 prefix |= new_bits;
468 bits_read += bit_count_to_read;479 bits_read += bit_count_to_read;
469 const result = try huffman_tree.query(huffman_tree_index, prefix);480 const result = huffman_tree.query(huffman_tree_index, prefix) catch |err| {
481 return err;
482 };
470483
471 switch (result) {484 switch (result) {
472 .symbol => |sym| {485 .symbol => |sym| {
...@@ -589,11 +602,14 @@ pub fn decodeZStandardFrame(...@@ -589,11 +602,14 @@ pub fn decodeZStandardFrame(
589 dest: []u8,602 dest: []u8,
590 src: []const u8,603 src: []const u8,
591 verify_checksum: bool,604 verify_checksum: bool,
592) (error{ UnknownContentSizeUnsupported, ContentTooLarge } || FrameError)!ReadWriteCount {605) (error{ UnknownContentSizeUnsupported, ContentTooLarge, EndOfStream } || FrameError)!ReadWriteCount {
593 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);606 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
594 var consumed_count: usize = 4;607 var consumed_count: usize = 4;
595608
596 const frame_header = try decodeZStandardHeader(src[consumed_count..], &consumed_count);609 var fbs = std.io.fixedBufferStream(src[consumed_count..]);
610 var source = fbs.reader();
611 const frame_header = try decodeZStandardHeader(source);
612 consumed_count += fbs.pos;
597613
598 if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported;614 if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported;
599615
...@@ -649,18 +665,25 @@ pub const FrameContext = struct {...@@ -649,18 +665,25 @@ pub const FrameContext = struct {
649/// `decodeZStandardFrame()`. Returns `error.WindowSizeUnknown` if the frame665/// `decodeZStandardFrame()`. Returns `error.WindowSizeUnknown` if the frame
650/// does not declare its content size or a window descriptor (this indicates a666/// does not declare its content size or a window descriptor (this indicates a
651/// malformed frame).667/// malformed frame).
668///
669/// Errors:
670/// - returns `error.WindowTooLarge`
671/// - returns `error.WindowSizeUnknown`
652pub fn decodeZStandardFrameAlloc(672pub fn decodeZStandardFrameAlloc(
653 allocator: std.mem.Allocator,673 allocator: std.mem.Allocator,
654 src: []const u8,674 src: []const u8,
655 verify_checksum: bool,675 verify_checksum: bool,
656 window_size_max: usize,676 window_size_max: usize,
657) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory } || FrameError)![]u8 {677) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory, EndOfStream } || FrameError)![]u8 {
658 var result = std.ArrayList(u8).init(allocator);678 var result = std.ArrayList(u8).init(allocator);
659 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);679 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
660 var consumed_count: usize = 4;680 var consumed_count: usize = 4;
661681
662 var frame_context = context: {682 var frame_context = context: {
663 const frame_header = try decodeZStandardHeader(src[consumed_count..], &consumed_count);683 var fbs = std.io.fixedBufferStream(src[consumed_count..]);
684 var source = fbs.reader();
685 const frame_header = try decodeZStandardHeader(source);
686 consumed_count += fbs.pos;
664 break :context try FrameContext.init(frame_header, window_size_max, verify_checksum);687 break :context try FrameContext.init(frame_header, window_size_max, verify_checksum);
665 };688 };
666689
...@@ -674,30 +697,7 @@ pub fn decodeZStandardFrameAlloc(...@@ -674,30 +697,7 @@ pub fn decodeZStandardFrameAlloc(
674697
675 var block_header = decodeBlockHeader(src[consumed_count..][0..3]);698 var block_header = decodeBlockHeader(src[consumed_count..][0..3]);
676 consumed_count += 3;699 consumed_count += 3;
677 var decode_state = DecodeState{700 var decode_state = DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
678 .repeat_offsets = .{
679 types.compressed_block.start_repeated_offset_1,
680 types.compressed_block.start_repeated_offset_2,
681 types.compressed_block.start_repeated_offset_3,
682 },
683
684 .offset = undefined,
685 .match = undefined,
686 .literal = undefined,
687
688 .literal_fse_buffer = &literal_fse_data,
689 .match_fse_buffer = &match_fse_data,
690 .offset_fse_buffer = &offset_fse_data,
691
692 .fse_tables_undefined = true,
693
694 .literal_written_count = 0,
695 .literal_header = undefined,
696 .literal_streams = undefined,
697 .literal_stream_reader = undefined,
698 .literal_stream_index = undefined,
699 .huffman_tree = null,
700 };
701 while (true) : ({701 while (true) : ({
702 block_header = decodeBlockHeader(src[consumed_count..][0..3]);702 block_header = decodeBlockHeader(src[consumed_count..][0..3]);
703 consumed_count += 3;703 consumed_count += 3;
...@@ -754,30 +754,7 @@ pub fn decodeFrameBlocks(...@@ -754,30 +754,7 @@ pub fn decodeFrameBlocks(
754 var block_header = decodeBlockHeader(src[0..3]);754 var block_header = decodeBlockHeader(src[0..3]);
755 var bytes_read: usize = 3;755 var bytes_read: usize = 3;
756 defer consumed_count.* += bytes_read;756 defer consumed_count.* += bytes_read;
757 var decode_state = DecodeState{757 var decode_state = DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
758 .repeat_offsets = .{
759 types.compressed_block.start_repeated_offset_1,
760 types.compressed_block.start_repeated_offset_2,
761 types.compressed_block.start_repeated_offset_3,
762 },
763
764 .offset = undefined,
765 .match = undefined,
766 .literal = undefined,
767
768 .literal_fse_buffer = &literal_fse_data,
769 .match_fse_buffer = &match_fse_data,
770 .offset_fse_buffer = &offset_fse_data,
771
772 .fse_tables_undefined = true,
773
774 .literal_written_count = 0,
775 .literal_header = undefined,
776 .literal_streams = undefined,
777 .literal_stream_reader = undefined,
778 .literal_stream_index = undefined,
779 .huffman_tree = null,
780 };
781 var written_count: usize = 0;758 var written_count: usize = 0;
782 while (true) : ({759 while (true) : ({
783 block_header = decodeBlockHeader(src[bytes_read..][0..3]);760 block_header = decodeBlockHeader(src[bytes_read..][0..3]);
...@@ -798,62 +775,6 @@ pub fn decodeFrameBlocks(...@@ -798,62 +775,6 @@ pub fn decodeFrameBlocks(
798 return written_count;775 return written_count;
799}776}
800777
801fn decodeRawBlock(
802 dest: []u8,
803 src: []const u8,
804 block_size: u21,
805 consumed_count: *usize,
806) error{MalformedBlockSize}!usize {
807 if (src.len < block_size) return error.MalformedBlockSize;
808 const data = src[0..block_size];
809 std.mem.copy(u8, dest, data);
810 consumed_count.* += block_size;
811 return block_size;
812}
813
814fn decodeRawBlockRingBuffer(
815 dest: *RingBuffer,
816 src: []const u8,
817 block_size: u21,
818 consumed_count: *usize,
819) error{MalformedBlockSize}!usize {
820 if (src.len < block_size) return error.MalformedBlockSize;
821 const data = src[0..block_size];
822 dest.writeSliceAssumeCapacity(data);
823 consumed_count.* += block_size;
824 return block_size;
825}
826
827fn decodeRleBlock(
828 dest: []u8,
829 src: []const u8,
830 block_size: u21,
831 consumed_count: *usize,
832) error{MalformedRleBlock}!usize {
833 if (src.len < 1) return error.MalformedRleBlock;
834 var write_pos: usize = 0;
835 while (write_pos < block_size) : (write_pos += 1) {
836 dest[write_pos] = src[0];
837 }
838 consumed_count.* += 1;
839 return block_size;
840}
841
842fn decodeRleBlockRingBuffer(
843 dest: *RingBuffer,
844 src: []const u8,
845 block_size: u21,
846 consumed_count: *usize,
847) error{MalformedRleBlock}!usize {
848 if (src.len < 1) return error.MalformedRleBlock;
849 var write_pos: usize = 0;
850 while (write_pos < block_size) : (write_pos += 1) {
851 dest.writeAssumeCapacity(src[0]);
852 }
853 consumed_count.* += 1;
854 return block_size;
855}
856
857/// Decode a single block from `src` into `dest`. The beginning of `src` should778/// Decode a single block from `src` into `dest`. The beginning of `src` should
858/// be the start of the block content (i.e. directly after the block header).779/// be the start of the block content (i.e. directly after the block header).
859/// Increments `consumed_count` by the number of bytes read from `src` to decode780/// Increments `consumed_count` by the number of bytes read from `src` to decode
...@@ -870,19 +791,37 @@ pub fn decodeBlock(...@@ -870,19 +791,37 @@ pub fn decodeBlock(
870 const block_size = block_header.block_size;791 const block_size = block_header.block_size;
871 if (block_size_max < block_size) return error.BlockSizeOverMaximum;792 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
872 switch (block_header.block_type) {793 switch (block_header.block_type) {
873 .raw => return decodeRawBlock(dest[written_count..], src, block_size, consumed_count),794 .raw => {
874 .rle => return decodeRleBlock(dest[written_count..], src, block_size, consumed_count),795 if (src.len < block_size) return error.MalformedBlockSize;
796 const data = src[0..block_size];
797 std.mem.copy(u8, dest[written_count..], data);
798 consumed_count.* += block_size;
799 return block_size;
800 },
801 .rle => {
802 if (src.len < 1) return error.MalformedRleBlock;
803 var write_pos: usize = written_count;
804 while (write_pos < block_size + written_count) : (write_pos += 1) {
805 dest[write_pos] = src[0];
806 }
807 consumed_count.* += 1;
808 return block_size;
809 },
875 .compressed => {810 .compressed => {
876 if (src.len < block_size) return error.MalformedBlockSize;811 if (src.len < block_size) return error.MalformedBlockSize;
877 var bytes_read: usize = 0;812 var bytes_read: usize = 0;
878 const literals = decodeLiteralsSection(src, &bytes_read) catch813 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
879 return error.MalformedCompressedBlock;814 return error.MalformedCompressedBlock;
880 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch815 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
816 const fbs_reader = fbs.reader();
817 const sequences_header = decodeSequencesHeader(fbs_reader) catch
881 return error.MalformedCompressedBlock;818 return error.MalformedCompressedBlock;
882819
883 bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch820 decode_state.prepare(fbs_reader, literals, sequences_header) catch
884 return error.MalformedCompressedBlock;821 return error.MalformedCompressedBlock;
885822
823 bytes_read += fbs.pos;
824
886 var bytes_written: usize = 0;825 var bytes_written: usize = 0;
887 if (sequences_header.sequence_count > 0) {826 if (sequences_header.sequence_count > 0) {
888 const bit_stream_bytes = src[bytes_read..block_size];827 const bit_stream_bytes = src[bytes_read..block_size];
...@@ -938,19 +877,37 @@ pub fn decodeBlockRingBuffer(...@@ -938,19 +877,37 @@ pub fn decodeBlockRingBuffer(
938 const block_size = block_header.block_size;877 const block_size = block_header.block_size;
939 if (block_size_max < block_size) return error.BlockSizeOverMaximum;878 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
940 switch (block_header.block_type) {879 switch (block_header.block_type) {
941 .raw => return decodeRawBlockRingBuffer(dest, src, block_size, consumed_count),880 .raw => {
942 .rle => return decodeRleBlockRingBuffer(dest, src, block_size, consumed_count),881 if (src.len < block_size) return error.MalformedBlockSize;
882 const data = src[0..block_size];
883 dest.writeSliceAssumeCapacity(data);
884 consumed_count.* += block_size;
885 return block_size;
886 },
887 .rle => {
888 if (src.len < 1) return error.MalformedRleBlock;
889 var write_pos: usize = 0;
890 while (write_pos < block_size) : (write_pos += 1) {
891 dest.writeAssumeCapacity(src[0]);
892 }
893 consumed_count.* += 1;
894 return block_size;
895 },
943 .compressed => {896 .compressed => {
944 if (src.len < block_size) return error.MalformedBlockSize;897 if (src.len < block_size) return error.MalformedBlockSize;
945 var bytes_read: usize = 0;898 var bytes_read: usize = 0;
946 const literals = decodeLiteralsSection(src, &bytes_read) catch899 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
947 return error.MalformedCompressedBlock;900 return error.MalformedCompressedBlock;
948 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch901 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
902 const fbs_reader = fbs.reader();
903 const sequences_header = decodeSequencesHeader(fbs_reader) catch
949 return error.MalformedCompressedBlock;904 return error.MalformedCompressedBlock;
950905
951 bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch906 decode_state.prepare(fbs_reader, literals, sequences_header) catch
952 return error.MalformedCompressedBlock;907 return error.MalformedCompressedBlock;
953908
909 bytes_read += fbs.pos;
910
954 var bytes_written: usize = 0;911 var bytes_written: usize = 0;
955 if (sequences_header.sequence_count > 0) {912 if (sequences_header.sequence_count > 0) {
956 const bit_stream_bytes = src[bytes_read..block_size];913 const bit_stream_bytes = src[bytes_read..block_size];
...@@ -991,6 +948,82 @@ pub fn decodeBlockRingBuffer(...@@ -991,6 +948,82 @@ pub fn decodeBlockRingBuffer(
991 }948 }
992}949}
993950
951/// Decode a single block from `source` into `dest`. Literal and sequence data
952/// from the block is copied into `literals_buffer` and `sequence_buffer`, which
953/// must be large enough or `error.LiteralsBufferTooSmall` and
954/// `error.SequenceBufferTooSmall` are returned (the maximum block size is an
955/// upper bound for the size of both buffers). See `decodeBlock`
956/// and `decodeBlockRingBuffer` for function that can decode a block without
957/// these extra copies.
958pub fn decodeBlockReader(
959 dest: *RingBuffer,
960 source: anytype,
961 block_header: frame.ZStandard.Block.Header,
962 decode_state: *DecodeState,
963 block_size_max: usize,
964 literals_buffer: []u8,
965 sequence_buffer: []u8,
966) !void {
967 const block_size = block_header.block_size;
968 var block_reader_limited = std.io.limitedReader(source, block_size);
969 const block_reader = block_reader_limited.reader();
970 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
971 switch (block_header.block_type) {
972 .raw => {
973 const slice = dest.sliceAt(dest.write_index, block_size);
974 try source.readNoEof(slice.first);
975 try source.readNoEof(slice.second);
976 dest.write_index = dest.mask2(dest.write_index + block_size);
977 },
978 .rle => {
979 const byte = try source.readByte();
980 var i: usize = 0;
981 while (i < block_size) : (i += 1) {
982 dest.writeAssumeCapacity(byte);
983 }
984 },
985 .compressed => {
986 const literals = try decodeLiteralsSection(block_reader, literals_buffer);
987 const sequences_header = try decodeSequencesHeader(block_reader);
988
989 try decode_state.prepare(block_reader, literals, sequences_header);
990
991 if (sequences_header.sequence_count > 0) {
992 if (sequence_buffer.len < block_reader_limited.bytes_left)
993 return error.SequenceBufferTooSmall;
994
995 const size = try block_reader.readAll(sequence_buffer);
996 var bit_stream: ReverseBitReader = undefined;
997 try bit_stream.init(sequence_buffer[0..size]);
998
999 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
1000
1001 var sequence_size_limit = block_size_max;
1002 var i: usize = 0;
1003 while (i < sequences_header.sequence_count) : (i += 1) {
1004 const decompressed_size = decode_state.decodeSequenceRingBuffer(
1005 dest,
1006 &bit_stream,
1007 sequence_size_limit,
1008 i == sequences_header.sequence_count - 1,
1009 ) catch return error.MalformedCompressedBlock;
1010 sequence_size_limit -= decompressed_size;
1011 }
1012 }
1013
1014 if (decode_state.literal_written_count < literals.header.regenerated_size) {
1015 const len = literals.header.regenerated_size - decode_state.literal_written_count;
1016 decode_state.decodeLiteralsRingBuffer(dest, len) catch
1017 return error.MalformedCompressedBlock;
1018 }
1019
1020 decode_state.literal_written_count = 0;
1021 assert(block_reader.readByte() == error.EndOfStream);
1022 },
1023 .reserved => return error.ReservedBlock,
1024 }
1025}
1026
994/// Decode the header of a skippable frame.1027/// Decode the header of a skippable frame.
995pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {1028pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {
996 const magic = readInt(u32, src[0..4]);1029 const magic = readInt(u32, src[0..4]);
...@@ -1002,13 +1035,6 @@ pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {...@@ -1002,13 +1035,6 @@ pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {
1002 };1035 };
1003}1036}
10041037
1005/// Returns the content size of a skippable frame.
1006pub fn skippableFrameSize(src: *const [8]u8) usize {
1007 assert(isSkippableMagic(readInt(u32, src[0..4])));
1008 const frame_size = readInt(u32, src[4..8]);
1009 return frame_size;
1010}
1011
1012/// Returns the window size required to decompress a frame, or `null` if it cannot be1038/// Returns the window size required to decompress a frame, or `null` if it cannot be
1013/// determined, which indicates a malformed frame header.1039/// determined, which indicates a malformed frame header.
1014pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 {1040pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 {
...@@ -1023,40 +1049,37 @@ pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 {...@@ -1023,40 +1049,37 @@ pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 {
1023}1049}
10241050
1025const InvalidBit = error{ UnusedBitSet, ReservedBitSet };1051const InvalidBit = error{ UnusedBitSet, ReservedBitSet };
1026/// Decode the header of a Zstandard frame. Returns `error.UnusedBitSet` or1052/// Decode the header of a Zstandard frame.
1027/// `error.ReservedBitSet` if the corresponding bits are sets.1053///
1028pub fn decodeZStandardHeader(src: []const u8, consumed_count: ?*usize) InvalidBit!frame.ZStandard.Header {1054/// Errors:
1029 const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, src[0]);1055/// - returns `error.UnusedBitSet` if the unused bits of the header are set
1056/// - returns `error.ReservedBitSet` if the reserved bits of the header are
1057/// set
1058pub fn decodeZStandardHeader(source: anytype) (error{EndOfStream} || InvalidBit)!frame.ZStandard.Header {
1059 const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, try source.readByte());
10301060
1031 if (descriptor.unused) return error.UnusedBitSet;1061 if (descriptor.unused) return error.UnusedBitSet;
1032 if (descriptor.reserved) return error.ReservedBitSet;1062 if (descriptor.reserved) return error.ReservedBitSet;
10331063
1034 var bytes_read_count: usize = 1;
1035
1036 var window_descriptor: ?u8 = null;1064 var window_descriptor: ?u8 = null;
1037 if (!descriptor.single_segment_flag) {1065 if (!descriptor.single_segment_flag) {
1038 window_descriptor = src[bytes_read_count];1066 window_descriptor = try source.readByte();
1039 bytes_read_count += 1;
1040 }1067 }
10411068
1042 var dictionary_id: ?u32 = null;1069 var dictionary_id: ?u32 = null;
1043 if (descriptor.dictionary_id_flag > 0) {1070 if (descriptor.dictionary_id_flag > 0) {
1044 // if flag is 3 then field_size = 4, else field_size = flag1071 // if flag is 3 then field_size = 4, else field_size = flag
1045 const field_size = (@as(u4, 1) << descriptor.dictionary_id_flag) >> 1;1072 const field_size = (@as(u4, 1) << descriptor.dictionary_id_flag) >> 1;
1046 dictionary_id = readVarInt(u32, src[bytes_read_count .. bytes_read_count + field_size]);1073 dictionary_id = try source.readVarInt(u32, .Little, field_size);
1047 bytes_read_count += field_size;
1048 }1074 }
10491075
1050 var content_size: ?u64 = null;1076 var content_size: ?u64 = null;
1051 if (descriptor.single_segment_flag or descriptor.content_size_flag > 0) {1077 if (descriptor.single_segment_flag or descriptor.content_size_flag > 0) {
1052 const field_size = @as(u4, 1) << descriptor.content_size_flag;1078 const field_size = @as(u4, 1) << descriptor.content_size_flag;
1053 content_size = readVarInt(u64, src[bytes_read_count .. bytes_read_count + field_size]);1079 content_size = try source.readVarInt(u64, .Little, field_size);
1054 if (field_size == 2) content_size.? += 256;1080 if (field_size == 2) content_size.? += 256;
1055 bytes_read_count += field_size;
1056 }1081 }
10571082
1058 if (consumed_count) |p| p.* += bytes_read_count;
1059
1060 const header = frame.ZStandard.Header{1083 const header = frame.ZStandard.Header{
1061 .descriptor = descriptor,1084 .descriptor = descriptor,
1062 .window_descriptor = window_descriptor,1085 .window_descriptor = window_descriptor,
...@@ -1080,12 +1103,20 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {...@@ -1080,12 +1103,20 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {
10801103
1081/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the1104/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
1082/// number of bytes the section uses.1105/// number of bytes the section uses.
1083pub fn decodeLiteralsSection(1106///
1107/// Errors:
1108/// - returns `error.MalformedLiteralsHeader` if the header is invalid
1109/// - returns `error.MalformedLiteralsSection` if there are errors decoding
1110pub fn decodeLiteralsSectionSlice(
1084 src: []const u8,1111 src: []const u8,
1085 consumed_count: *usize,1112 consumed_count: *usize,
1086) (error{ MalformedLiteralsHeader, MalformedLiteralsSection } || DecodeHuffmanError)!LiteralsSection {1113) (error{ MalformedLiteralsHeader, MalformedLiteralsSection, EndOfStream } || DecodeHuffmanError)!LiteralsSection {
1087 var bytes_read: usize = 0;1114 var bytes_read: usize = 0;
1088 const header = try decodeLiteralsHeader(src, &bytes_read);1115 const header = header: {
1116 var fbs = std.io.fixedBufferStream(src);
1117 defer bytes_read = fbs.pos;
1118 break :header decodeLiteralsHeader(fbs.reader()) catch return error.MalformedLiteralsHeader;
1119 };
1089 switch (header.block_type) {1120 switch (header.block_type) {
1090 .raw => {1121 .raw => {
1091 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;1122 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;
...@@ -1110,7 +1141,7 @@ pub fn decodeLiteralsSection(...@@ -1110,7 +1141,7 @@ pub fn decodeLiteralsSection(
1110 .compressed, .treeless => {1141 .compressed, .treeless => {
1111 const huffman_tree_start = bytes_read;1142 const huffman_tree_start = bytes_read;
1112 const huffman_tree = if (header.block_type == .compressed)1143 const huffman_tree = if (header.block_type == .compressed)
1113 try decodeHuffmanTree(src[bytes_read..], &bytes_read)1144 try decodeHuffmanTreeSlice(src[bytes_read..], &bytes_read)
1114 else1145 else
1115 null;1146 null;
1116 const huffman_tree_size = bytes_read - huffman_tree_start;1147 const huffman_tree_size = bytes_read - huffman_tree_start;
...@@ -1119,137 +1150,185 @@ pub fn decodeLiteralsSection(...@@ -1119,137 +1150,185 @@ pub fn decodeLiteralsSection(
1119 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;1150 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;
1120 const stream_data = src[bytes_read .. bytes_read + total_streams_size];1151 const stream_data = src[bytes_read .. bytes_read + total_streams_size];
11211152
1122 if (header.size_format == 0) {1153 const streams = try decodeStreams(header.size_format, stream_data);
1123 consumed_count.* += total_streams_size + bytes_read;1154 consumed_count.* += bytes_read + total_streams_size;
1124 return LiteralsSection{1155 return LiteralsSection{
1125 .header = header,1156 .header = header,
1126 .huffman_tree = huffman_tree,1157 .huffman_tree = huffman_tree,
1127 .streams = .{ .one = stream_data },1158 .streams = streams,
1128 };1159 };
1129 }1160 },
11301161 }
1131 if (stream_data.len < 6) return error.MalformedLiteralsSection;1162}
1132
1133 const stream_1_length = @as(usize, readInt(u16, stream_data[0..2]));
1134 const stream_2_length = @as(usize, readInt(u16, stream_data[2..4]));
1135 const stream_3_length = @as(usize, readInt(u16, stream_data[4..6]));
1136 const stream_4_length = (total_streams_size - 6) - (stream_1_length + stream_2_length + stream_3_length);
11371163
1138 const stream_1_start = 6;1164/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
1139 const stream_2_start = stream_1_start + stream_1_length;1165/// number of bytes the section uses.
1140 const stream_3_start = stream_2_start + stream_2_length;1166///
1141 const stream_4_start = stream_3_start + stream_3_length;1167/// Errors:
1168/// - returns `error.MalformedLiteralsHeader` if the header is invalid
1169/// - returns `error.MalformedLiteralsSection` if there are errors decoding
1170pub fn decodeLiteralsSection(
1171 source: anytype,
1172 buffer: []u8,
1173) !LiteralsSection {
1174 const header = try decodeLiteralsHeader(source);
1175 switch (header.block_type) {
1176 .raw => {
1177 try source.readNoEof(buffer[0..header.regenerated_size]);
1178 return LiteralsSection{
1179 .header = header,
1180 .huffman_tree = null,
1181 .streams = .{ .one = buffer },
1182 };
1183 },
1184 .rle => {
1185 buffer[0] = try source.readByte();
1186 return LiteralsSection{
1187 .header = header,
1188 .huffman_tree = null,
1189 .streams = .{ .one = buffer[0..1] },
1190 };
1191 },
1192 .compressed, .treeless => {
1193 var counting_reader = std.io.countingReader(source);
1194 const huffman_tree = if (header.block_type == .compressed)
1195 try decodeHuffmanTree(counting_reader.reader(), buffer)
1196 else
1197 null;
1198 const huffman_tree_size = counting_reader.bytes_read;
1199 const total_streams_size = @as(usize, header.compressed_size.?) - @intCast(usize, huffman_tree_size);
11421200
1143 if (stream_data.len < stream_4_start + stream_4_length) return error.MalformedLiteralsSection;1201 if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall;
1144 consumed_count.* += total_streams_size + bytes_read;1202 try source.readNoEof(buffer[0..total_streams_size]);
1203 const stream_data = buffer[0..total_streams_size];
11451204
1205 const streams = try decodeStreams(header.size_format, stream_data);
1146 return LiteralsSection{1206 return LiteralsSection{
1147 .header = header,1207 .header = header,
1148 .huffman_tree = huffman_tree,1208 .huffman_tree = huffman_tree,
1149 .streams = .{ .four = .{1209 .streams = streams,
1150 stream_data[stream_1_start .. stream_1_start + stream_1_length],
1151 stream_data[stream_2_start .. stream_2_start + stream_2_length],
1152 stream_data[stream_3_start .. stream_3_start + stream_3_length],
1153 stream_data[stream_4_start .. stream_4_start + stream_4_length],
1154 } },
1155 };1210 };
1156 },1211 },
1157 }1212 }
1158}1213}
11591214
1215fn decodeStreams(size_format: u2, stream_data: []const u8) !LiteralsSection.Streams {
1216 if (size_format == 0) {
1217 return .{ .one = stream_data };
1218 }
1219
1220 if (stream_data.len < 6) return error.MalformedLiteralsSection;
1221
1222 const stream_1_length = @as(usize, readInt(u16, stream_data[0..2]));
1223 const stream_2_length = @as(usize, readInt(u16, stream_data[2..4]));
1224 const stream_3_length = @as(usize, readInt(u16, stream_data[4..6]));
1225
1226 const stream_1_start = 6;
1227 const stream_2_start = stream_1_start + stream_1_length;
1228 const stream_3_start = stream_2_start + stream_2_length;
1229 const stream_4_start = stream_3_start + stream_3_length;
1230
1231 return .{ .four = .{
1232 stream_data[stream_1_start .. stream_1_start + stream_1_length],
1233 stream_data[stream_2_start .. stream_2_start + stream_2_length],
1234 stream_data[stream_3_start .. stream_3_start + stream_3_length],
1235 stream_data[stream_4_start..],
1236 } };
1237}
1238
1160const DecodeHuffmanError = error{1239const DecodeHuffmanError = error{
1161 MalformedHuffmanTree,1240 MalformedHuffmanTree,
1162 MalformedFseTable,1241 MalformedFseTable,
1163 MalformedAccuracyLog,1242 MalformedAccuracyLog,
1164};1243};
11651244
1166fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError!LiteralsSection.HuffmanTree {1245fn decodeFseHuffmanTree(source: anytype, compressed_size: usize, buffer: []u8, weights: *[256]u4) !usize {
1167 var bytes_read: usize = 0;1246 var stream = std.io.limitedReader(source, compressed_size);
1168 bytes_read += 1;1247 var bit_reader = bitReader(stream.reader());
1169 if (src.len == 0) return error.MalformedHuffmanTree;
1170 const header = src[0];
1171 var symbol_count: usize = undefined;
1172 var weights: [256]u4 = undefined;
1173 var max_number_of_bits: u4 = undefined;
1174 if (header < 128) {
1175 // FSE compressed weights
1176 const compressed_size = header;
1177 if (src.len < 1 + compressed_size) return error.MalformedHuffmanTree;
1178 var stream = std.io.fixedBufferStream(src[1 .. compressed_size + 1]);
1179 var counting_reader = std.io.countingReader(stream.reader());
1180 var bit_reader = bitReader(counting_reader.reader());
1181
1182 var entries: [1 << 6]Table.Fse = undefined;
1183 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
1184 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
1185 error.EndOfStream => return error.MalformedFseTable,
1186 };
1187 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
1188
1189 const start_index = std.math.cast(usize, 1 + counting_reader.bytes_read) orelse return error.MalformedHuffmanTree;
1190 var huff_data = src[start_index .. compressed_size + 1];
1191 var huff_bits: ReverseBitReader = undefined;
1192 huff_bits.init(huff_data) catch return error.MalformedHuffmanTree;
1193
1194 var i: usize = 0;
1195 var even_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
1196 var odd_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
1197
1198 while (i < 255) {
1199 const even_data = entries[even_state];
1200 var read_bits: usize = 0;
1201 const even_bits = huff_bits.readBits(u32, even_data.bits, &read_bits) catch unreachable;
1202 weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree;
1203 i += 1;
1204 if (read_bits < even_data.bits) {
1205 weights[i] = std.math.cast(u4, entries[odd_state].symbol) orelse return error.MalformedHuffmanTree;
1206 i += 1;
1207 break;
1208 }
1209 even_state = even_data.baseline + even_bits;
12101248
1211 read_bits = 0;1249 var entries: [1 << 6]Table.Fse = undefined;
1212 const odd_data = entries[odd_state];1250 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
1213 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;1251 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
1214 weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree;1252 error.EndOfStream => return error.MalformedFseTable,
1215 i += 1;1253 };
1216 if (read_bits < odd_data.bits) {1254 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
1217 if (i == 256) return error.MalformedHuffmanTree;
1218 weights[i] = std.math.cast(u4, entries[even_state].symbol) orelse return error.MalformedHuffmanTree;
1219 i += 1;
1220 break;
1221 }
1222 odd_state = odd_data.baseline + odd_bits;
1223 } else return error.MalformedHuffmanTree;
12241255
1225 symbol_count = i + 1; // stream contains all but the last symbol1256 const amount = try stream.reader().readAll(buffer);
1226 bytes_read += compressed_size;1257 var huff_bits: ReverseBitReader = undefined;
1227 } else {1258 huff_bits.init(buffer[0..amount]) catch return error.MalformedHuffmanTree;
1228 const encoded_symbol_count = header - 127;1259
1229 symbol_count = encoded_symbol_count + 1;1260 return assignWeights(&huff_bits, accuracy_log, &entries, weights);
1230 const weights_byte_count = (encoded_symbol_count + 1) / 2;1261}
1231 if (src.len < weights_byte_count) return error.MalformedHuffmanTree;1262
1232 var i: usize = 0;1263fn decodeFseHuffmanTreeSlice(src: []const u8, compressed_size: usize, weights: *[256]u4) !usize {
1233 while (i < weights_byte_count) : (i += 1) {1264 if (src.len < compressed_size) return error.MalformedHuffmanTree;
1234 weights[2 * i] = @intCast(u4, src[i + 1] >> 4);1265 var stream = std.io.fixedBufferStream(src[0..compressed_size]);
1235 weights[2 * i + 1] = @intCast(u4, src[i + 1] & 0xF);1266 var counting_reader = std.io.countingReader(stream.reader());
1267 var bit_reader = bitReader(counting_reader.reader());
1268
1269 var entries: [1 << 6]Table.Fse = undefined;
1270 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
1271 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
1272 error.EndOfStream => return error.MalformedFseTable,
1273 };
1274 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
1275
1276 const start_index = std.math.cast(usize, counting_reader.bytes_read) orelse return error.MalformedHuffmanTree;
1277 var huff_data = src[start_index..compressed_size];
1278 var huff_bits: ReverseBitReader = undefined;
1279 huff_bits.init(huff_data) catch return error.MalformedHuffmanTree;
1280
1281 return assignWeights(&huff_bits, accuracy_log, &entries, weights);
1282}
1283
1284fn assignWeights(huff_bits: *ReverseBitReader, accuracy_log: usize, entries: *[1 << 6]Table.Fse, weights: *[256]u4) !usize {
1285 var i: usize = 0;
1286 var even_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
1287 var odd_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
1288
1289 while (i < 255) {
1290 const even_data = entries[even_state];
1291 var read_bits: usize = 0;
1292 const even_bits = huff_bits.readBits(u32, even_data.bits, &read_bits) catch unreachable;
1293 weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree;
1294 i += 1;
1295 if (read_bits < even_data.bits) {
1296 weights[i] = std.math.cast(u4, entries[odd_state].symbol) orelse return error.MalformedHuffmanTree;
1297 i += 1;
1298 break;
1236 }1299 }
1237 bytes_read += weights_byte_count;1300 even_state = even_data.baseline + even_bits;
1238 }1301
1239 var weight_power_sum: u16 = 0;1302 read_bits = 0;
1240 for (weights[0 .. symbol_count - 1]) |value| {1303 const odd_data = entries[odd_state];
1241 if (value > 0) {1304 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;
1242 weight_power_sum += @as(u16, 1) << (value - 1);1305 weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree;
1306 i += 1;
1307 if (read_bits < odd_data.bits) {
1308 if (i == 256) return error.MalformedHuffmanTree;
1309 weights[i] = std.math.cast(u4, entries[even_state].symbol) orelse return error.MalformedHuffmanTree;
1310 i += 1;
1311 break;
1243 }1312 }
1244 }1313 odd_state = odd_data.baseline + odd_bits;
1314 } else return error.MalformedHuffmanTree;
12451315
1246 // advance to next power of two (even if weight_power_sum is a power of 2)1316 return i + 1; // stream contains all but the last symbol
1247 max_number_of_bits = std.math.log2_int(u16, weight_power_sum) + 1;1317}
1248 const next_power_of_two = @as(u16, 1) << max_number_of_bits;
1249 weights[symbol_count - 1] = std.math.log2_int(u16, next_power_of_two - weight_power_sum) + 1;
12501318
1251 var weight_sorted_prefixed_symbols: [256]LiteralsSection.HuffmanTree.PrefixedSymbol = undefined;1319fn decodeDirectHuffmanTree(source: anytype, encoded_symbol_count: usize, weights: *[256]u4) !usize {
1252 for (weight_sorted_prefixed_symbols[0..symbol_count]) |_, i| {1320 const weights_byte_count = (encoded_symbol_count + 1) / 2;
1321 var i: usize = 0;
1322 while (i < weights_byte_count) : (i += 1) {
1323 const byte = try source.readByte();
1324 weights[2 * i] = @intCast(u4, byte >> 4);
1325 weights[2 * i + 1] = @intCast(u4, byte & 0xF);
1326 }
1327 return encoded_symbol_count + 1;
1328}
1329
1330fn assignSymbols(weight_sorted_prefixed_symbols: []LiteralsSection.HuffmanTree.PrefixedSymbol, weights: [256]u4) usize {
1331 for (weight_sorted_prefixed_symbols) |_, i| {
1253 weight_sorted_prefixed_symbols[i] = .{1332 weight_sorted_prefixed_symbols[i] = .{
1254 .symbol = @intCast(u8, i),1333 .symbol = @intCast(u8, i),
1255 .weight = undefined,1334 .weight = undefined,
...@@ -1259,7 +1338,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError...@@ -1259,7 +1338,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
12591338
1260 std.sort.sort(1339 std.sort.sort(
1261 LiteralsSection.HuffmanTree.PrefixedSymbol,1340 LiteralsSection.HuffmanTree.PrefixedSymbol,
1262 weight_sorted_prefixed_symbols[0..symbol_count],1341 weight_sorted_prefixed_symbols,
1263 weights,1342 weights,
1264 lessThanByWeight,1343 lessThanByWeight,
1265 );1344 );
...@@ -1267,6 +1346,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError...@@ -1267,6 +1346,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
1267 var prefix: u16 = 0;1346 var prefix: u16 = 0;
1268 var prefixed_symbol_count: usize = 0;1347 var prefixed_symbol_count: usize = 0;
1269 var sorted_index: usize = 0;1348 var sorted_index: usize = 0;
1349 const symbol_count = weight_sorted_prefixed_symbols.len;
1270 while (sorted_index < symbol_count) {1350 while (sorted_index < symbol_count) {
1271 var symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;1351 var symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
1272 const weight = weights[symbol];1352 const weight = weights[symbol];
...@@ -1290,7 +1370,24 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError...@@ -1290,7 +1370,24 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
1290 weight_sorted_prefixed_symbols[prefixed_symbol_count].weight = weight;1370 weight_sorted_prefixed_symbols[prefixed_symbol_count].weight = weight;
1291 }1371 }
1292 }1372 }
1293 consumed_count.* += bytes_read;1373 return prefixed_symbol_count;
1374}
1375
1376fn buildHuffmanTree(weights: *[256]u4, symbol_count: usize) LiteralsSection.HuffmanTree {
1377 var weight_power_sum: u16 = 0;
1378 for (weights[0 .. symbol_count - 1]) |value| {
1379 if (value > 0) {
1380 weight_power_sum += @as(u16, 1) << (value - 1);
1381 }
1382 }
1383
1384 // advance to next power of two (even if weight_power_sum is a power of 2)
1385 const max_number_of_bits = std.math.log2_int(u16, weight_power_sum) + 1;
1386 const next_power_of_two = @as(u16, 1) << max_number_of_bits;
1387 weights[symbol_count - 1] = std.math.log2_int(u16, next_power_of_two - weight_power_sum) + 1;
1388
1389 var weight_sorted_prefixed_symbols: [256]LiteralsSection.HuffmanTree.PrefixedSymbol = undefined;
1390 const prefixed_symbol_count = assignSymbols(weight_sorted_prefixed_symbols[0..symbol_count], weights.*);
1294 const tree = LiteralsSection.HuffmanTree{1391 const tree = LiteralsSection.HuffmanTree{
1295 .max_bit_count = max_number_of_bits,1392 .max_bit_count = max_number_of_bits,
1296 .symbol_count_minus_one = @intCast(u8, prefixed_symbol_count - 1),1393 .symbol_count_minus_one = @intCast(u8, prefixed_symbol_count - 1),
...@@ -1299,6 +1396,37 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError...@@ -1299,6 +1396,37 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
1299 return tree;1396 return tree;
1300}1397}
13011398
1399fn decodeHuffmanTree(source: anytype, buffer: []u8) !LiteralsSection.HuffmanTree {
1400 const header = try source.readByte();
1401 var weights: [256]u4 = undefined;
1402 const symbol_count = if (header < 128)
1403 // FSE compressed weights
1404 try decodeFseHuffmanTree(source, header, buffer, &weights)
1405 else
1406 try decodeDirectHuffmanTree(source, header - 127, &weights);
1407
1408 return buildHuffmanTree(&weights, symbol_count);
1409}
1410
1411fn decodeHuffmanTreeSlice(src: []const u8, consumed_count: *usize) (error{EndOfStream} || DecodeHuffmanError)!LiteralsSection.HuffmanTree {
1412 if (src.len == 0) return error.MalformedHuffmanTree;
1413 const header = src[0];
1414 var bytes_read: usize = 1;
1415 var weights: [256]u4 = undefined;
1416 const symbol_count = if (header < 128) count: {
1417 // FSE compressed weights
1418 bytes_read += header;
1419 break :count try decodeFseHuffmanTreeSlice(src[1..], header, &weights);
1420 } else count: {
1421 var fbs = std.io.fixedBufferStream(src[1..]);
1422 defer bytes_read += fbs.pos;
1423 break :count try decodeDirectHuffmanTree(fbs.reader(), header - 127, &weights);
1424 };
1425
1426 consumed_count.* += bytes_read;
1427 return buildHuffmanTree(&weights, symbol_count);
1428}
1429
1302fn lessThanByWeight(1430fn lessThanByWeight(
1303 weights: [256]u4,1431 weights: [256]u4,
1304 lhs: LiteralsSection.HuffmanTree.PrefixedSymbol,1432 lhs: LiteralsSection.HuffmanTree.PrefixedSymbol,
...@@ -1311,9 +1439,8 @@ fn lessThanByWeight(...@@ -1311,9 +1439,8 @@ fn lessThanByWeight(
1311}1439}
13121440
1313/// Decode a literals section header.1441/// Decode a literals section header.
1314pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{MalformedLiteralsHeader}!LiteralsSection.Header {1442pub fn decodeLiteralsHeader(source: anytype) !LiteralsSection.Header {
1315 if (src.len == 0) return error.MalformedLiteralsHeader;1443 const byte0 = try source.readByte();
1316 const byte0 = src[0];
1317 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);1444 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);
1318 const size_format = @intCast(u2, (byte0 & 0b1100) >> 2);1445 const size_format = @intCast(u2, (byte0 & 0b1100) >> 2);
1319 var regenerated_size: u20 = undefined;1446 var regenerated_size: u20 = undefined;
...@@ -1323,47 +1450,31 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{Malfo...@@ -1323,47 +1450,31 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{Malfo
1323 switch (size_format) {1450 switch (size_format) {
1324 0, 2 => {1451 0, 2 => {
1325 regenerated_size = byte0 >> 3;1452 regenerated_size = byte0 >> 3;
1326 consumed_count.* += 1;
1327 },
1328 1 => {
1329 if (src.len < 2) return error.MalformedLiteralsHeader;
1330 regenerated_size = (byte0 >> 4) +
1331 (@as(u20, src[1]) << 4);
1332 consumed_count.* += 2;
1333 },
1334 3 => {
1335 if (src.len < 3) return error.MalformedLiteralsHeader;
1336 regenerated_size = (byte0 >> 4) +
1337 (@as(u20, src[1]) << 4) +
1338 (@as(u20, src[2]) << 12);
1339 consumed_count.* += 3;
1340 },1453 },
1454 1 => regenerated_size = (byte0 >> 4) + (@as(u20, try source.readByte()) << 4),
1455 3 => regenerated_size = (byte0 >> 4) +
1456 (@as(u20, try source.readByte()) << 4) +
1457 (@as(u20, try source.readByte()) << 12),
1341 }1458 }
1342 },1459 },
1343 .compressed, .treeless => {1460 .compressed, .treeless => {
1344 const byte1 = src[1];1461 const byte1 = try source.readByte();
1345 const byte2 = src[2];1462 const byte2 = try source.readByte();
1346 switch (size_format) {1463 switch (size_format) {
1347 0, 1 => {1464 0, 1 => {
1348 if (src.len < 3) return error.MalformedLiteralsHeader;
1349 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);1465 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);
1350 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);1466 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);
1351 consumed_count.* += 3;
1352 },1467 },
1353 2 => {1468 2 => {
1354 if (src.len < 4) return error.MalformedLiteralsHeader;1469 const byte3 = try source.readByte();
1355 const byte3 = src[3];
1356 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);1470 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);
1357 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);1471 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);
1358 consumed_count.* += 4;
1359 },1472 },
1360 3 => {1473 3 => {
1361 if (src.len < 5) return error.MalformedLiteralsHeader;1474 const byte3 = try source.readByte();
1362 const byte3 = src[3];1475 const byte4 = try source.readByte();
1363 const byte4 = src[4];
1364 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);1476 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);
1365 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);1477 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);
1366 consumed_count.* += 5;
1367 },1478 },
1368 }1479 }
1369 },1480 },
...@@ -1377,18 +1488,17 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{Malfo...@@ -1377,18 +1488,17 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{Malfo
1377}1488}
13781489
1379/// Decode a sequences section header.1490/// Decode a sequences section header.
1491///
1492/// Errors:
1493/// - returns `error.ReservedBitSet` is the reserved bit is set
1494/// - returns `error.MalformedSequencesHeader` if the header is invalid
1380pub fn decodeSequencesHeader(1495pub fn decodeSequencesHeader(
1381 src: []const u8,1496 source: anytype,
1382 consumed_count: *usize,1497) !SequencesSection.Header {
1383) error{ MalformedSequencesHeader, ReservedBitSet }!SequencesSection.Header {
1384 if (src.len == 0) return error.MalformedSequencesHeader;
1385 var sequence_count: u24 = undefined;1498 var sequence_count: u24 = undefined;
13861499
1387 var bytes_read: usize = 0;1500 const byte0 = try source.readByte();
1388 const byte0 = src[0];
1389 if (byte0 == 0) {1501 if (byte0 == 0) {
1390 bytes_read += 1;
1391 consumed_count.* += bytes_read;
1392 return SequencesSection.Header{1502 return SequencesSection.Header{
1393 .sequence_count = 0,1503 .sequence_count = 0,
1394 .offsets = undefined,1504 .offsets = undefined,
...@@ -1397,22 +1507,14 @@ pub fn decodeSequencesHeader(...@@ -1397,22 +1507,14 @@ pub fn decodeSequencesHeader(
1397 };1507 };
1398 } else if (byte0 < 128) {1508 } else if (byte0 < 128) {
1399 sequence_count = byte0;1509 sequence_count = byte0;
1400 bytes_read += 1;
1401 } else if (byte0 < 255) {1510 } else if (byte0 < 255) {
1402 if (src.len < 2) return error.MalformedSequencesHeader;1511 sequence_count = (@as(u24, (byte0 - 128)) << 8) + try source.readByte();
1403 sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1];
1404 bytes_read += 2;
1405 } else {1512 } else {
1406 if (src.len < 3) return error.MalformedSequencesHeader;1513 sequence_count = (try source.readByte()) + (@as(u24, try source.readByte()) << 8) + 0x7F00;
1407 sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00;
1408 bytes_read += 3;
1409 }1514 }
14101515
1411 if (src.len < bytes_read + 1) return error.MalformedSequencesHeader;1516 const compression_modes = try source.readByte();
1412 const compression_modes = src[bytes_read];
1413 bytes_read += 1;
14141517
1415 consumed_count.* += bytes_read;
1416 const matches_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00001100) >> 2);1518 const matches_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00001100) >> 2);
1417 const offsets_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00110000) >> 4);1519 const offsets_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00110000) >> 4);
1418 const literal_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b11000000) >> 6);1520 const literal_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b11000000) >> 6);
...@@ -1615,7 +1717,7 @@ fn BitReader(comptime Reader: type) type {...@@ -1615,7 +1717,7 @@ fn BitReader(comptime Reader: type) type {
1615 };1717 };
1616}1718}
16171719
1618fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) {1720pub fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) {
1619 return .{ .underlying = std.io.bitReader(.Little, reader) };1721 return .{ .underlying = std.io.bitReader(.Little, reader) };
1620}1722}
16211723