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 {
1414 return std.mem.readVarInt(T, bytes, .Little);
1515}
1616
17fn isSkippableMagic(magic: u32) bool {
17pub fn isSkippableMagic(magic: u32) bool {
1818 return frame.Skippable.magic_number_min <= magic and magic <= frame.Skippable.magic_number_max;
1919}
2020
21/// Returns the decompressed size of the frame at the start of `src`. Returns 0
22/// if the the frame is skippable, `null` for Zstanndard frames that do not
23/// declare their content size. Returns `UnusedBitSet` and `ReservedBitSet`
24/// errors if the respective bits of the the frame descriptor are set.
25pub fn getFrameDecompressedSize(src: []const u8) (InvalidBit || error{BadMagic})!?u64 {
26 switch (try frameType(src)) {
27 .zstandard => {
28 const header = try decodeZStandardHeader(src[4..], null);
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]);
21/// Returns the kind of frame at the beginning of `src`.
22///
23/// Errors:
24/// - returns `error.BadMagic` if `source` begins with bytes not equal to the
25/// Zstandard frame magic number, or outside the range of magic numbers for
26/// skippable frames.
27pub fn decodeFrameType(source: anytype) !frame.Kind {
28 const magic = try source.readIntLittle(u32);
4029 return if (magic == frame.ZStandard.magic_number)
4130 .zstandard
4231 else if (isSkippableMagic(magic))
......@@ -52,15 +41,21 @@ const ReadWriteCount = struct {
5241
5342/// Decodes the frame at the start of `src` into `dest`. Returns the number of
5443/// bytes read from `src` and written to `dest`.
44///
45/// Errors:
46/// - returns `error.UnknownContentSizeUnsupported`
47/// - returns `error.ContentTooLarge`
48/// - returns `error.BadMagic`
5549pub fn decodeFrame(
5650 dest: []u8,
5751 src: []const u8,
5852 verify_checksum: bool,
59) (error{ UnknownContentSizeUnsupported, ContentTooLarge, BadMagic } || FrameError)!ReadWriteCount {
60 return switch (try frameType(src)) {
53) !ReadWriteCount {
54 var fbs = std.io.fixedBufferStream(src);
55 return switch (try decodeFrameType(fbs.reader())) {
6156 .zstandard => decodeZStandardFrame(dest, src, verify_checksum),
6257 .skippable => ReadWriteCount{
63 .read_count = skippableFrameSize(src[0..8]) + 8,
58 .read_count = try fbs.reader().readIntLittle(u32) + 8,
6459 .write_count = 0,
6560 },
6661 };
......@@ -97,16 +92,52 @@ pub const DecodeState = struct {
9792 };
9893 }
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
100126 /// 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`.
102 /// Returns `error.BitStreamHasNoStartBit` if the (reversed) literal bitstream's
103 /// first byte does not have any bits set.
127 /// stream and Huffman tree from `literals` and reads the FSE tables from
128 /// `source`.
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.
104135 pub fn prepare(
105136 self: *DecodeState,
106 src: []const u8,
137 source: anytype,
107138 literals: LiteralsSection,
108139 sequences_header: SequencesSection.Header,
109 ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize {
140 ) !void {
110141 self.literal_written_count = 0;
111142 self.literal_header = literals.header;
112143 self.literal_streams = literals.streams;
......@@ -129,28 +160,11 @@ pub const DecodeState = struct {
129160 }
130161
131162 if (sequences_header.sequence_count > 0) {
132 var bytes_read = try self.updateFseTable(
133 src,
134 .literal,
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 );
163 try self.updateFseTable(source, .literal, sequences_header.literal_lengths);
164 try self.updateFseTable(source, .offset, sequences_header.offsets);
165 try self.updateFseTable(source, .match, sequences_header.match_lengths);
149166 self.fse_tables_undefined = false;
150
151 return bytes_read;
152167 }
153 return 0;
154168 }
155169
156170 /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream`
......@@ -208,10 +222,10 @@ pub const DecodeState = struct {
208222
209223 fn updateFseTable(
210224 self: *DecodeState,
211 src: []const u8,
225 source: anytype,
212226 comptime choice: DataType,
213227 mode: SequencesSection.Header.Mode,
214 ) FseTableError!usize {
228 ) !void {
215229 const field_name = @tagName(choice);
216230 switch (mode) {
217231 .predefined => {
......@@ -220,17 +234,13 @@ pub const DecodeState = struct {
220234
221235 @field(self, field_name).table =
222236 @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table");
223 return 0;
224237 },
225238 .rle => {
226239 @field(self, field_name).accuracy_log = 0;
227 @field(self, field_name).table = .{ .rle = src[0] };
228 return 1;
240 @field(self, field_name).table = .{ .rle = try source.readByte() };
229241 },
230242 .fse => {
231 var stream = std.io.fixedBufferStream(src);
232 var counting_reader = std.io.countingReader(stream.reader());
233 var bit_reader = bitReader(counting_reader.reader());
243 var bit_reader = bitReader(source);
234244
235245 const table_size = try decodeFseTable(
236246 &bit_reader,
......@@ -242,9 +252,8 @@ pub const DecodeState = struct {
242252 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],
243253 };
244254 @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;
246255 },
247 .repeat => return if (self.fse_tables_undefined) error.RepeatModeFirst else 0,
256 .repeat => if (self.fse_tables_undefined) return error.RepeatModeFirst,
248257 }
249258 }
250259
......@@ -462,11 +471,15 @@ pub const DecodeState = struct {
462471 while (i < len) : (i += 1) {
463472 var prefix: u16 = 0;
464473 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 };
466477 prefix <<= bit_count_to_read;
467478 prefix |= new_bits;
468479 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
471484 switch (result) {
472485 .symbol => |sym| {
......@@ -589,11 +602,14 @@ pub fn decodeZStandardFrame(
589602 dest: []u8,
590603 src: []const u8,
591604 verify_checksum: bool,
592) (error{ UnknownContentSizeUnsupported, ContentTooLarge } || FrameError)!ReadWriteCount {
605) (error{ UnknownContentSizeUnsupported, ContentTooLarge, EndOfStream } || FrameError)!ReadWriteCount {
593606 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
594607 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
598614 if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported;
599615
......@@ -649,18 +665,25 @@ pub const FrameContext = struct {
649665/// `decodeZStandardFrame()`. Returns `error.WindowSizeUnknown` if the frame
650666/// does not declare its content size or a window descriptor (this indicates a
651667/// malformed frame).
668///
669/// Errors:
670/// - returns `error.WindowTooLarge`
671/// - returns `error.WindowSizeUnknown`
652672pub fn decodeZStandardFrameAlloc(
653673 allocator: std.mem.Allocator,
654674 src: []const u8,
655675 verify_checksum: bool,
656676 window_size_max: usize,
657) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory } || FrameError)![]u8 {
677) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory, EndOfStream } || FrameError)![]u8 {
658678 var result = std.ArrayList(u8).init(allocator);
659679 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
660680 var consumed_count: usize = 4;
661681
662682 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;
664687 break :context try FrameContext.init(frame_header, window_size_max, verify_checksum);
665688 };
666689
......@@ -674,30 +697,7 @@ pub fn decodeZStandardFrameAlloc(
674697
675698 var block_header = decodeBlockHeader(src[consumed_count..][0..3]);
676699 consumed_count += 3;
677 var decode_state = DecodeState{
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 };
700 var decode_state = DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
701701 while (true) : ({
702702 block_header = decodeBlockHeader(src[consumed_count..][0..3]);
703703 consumed_count += 3;
......@@ -754,30 +754,7 @@ pub fn decodeFrameBlocks(
754754 var block_header = decodeBlockHeader(src[0..3]);
755755 var bytes_read: usize = 3;
756756 defer consumed_count.* += bytes_read;
757 var decode_state = DecodeState{
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 };
757 var decode_state = DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
781758 var written_count: usize = 0;
782759 while (true) : ({
783760 block_header = decodeBlockHeader(src[bytes_read..][0..3]);
......@@ -798,62 +775,6 @@ pub fn decodeFrameBlocks(
798775 return written_count;
799776}
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
857778/// Decode a single block from `src` into `dest`. The beginning of `src` should
858779/// be the start of the block content (i.e. directly after the block header).
859780/// Increments `consumed_count` by the number of bytes read from `src` to decode
......@@ -870,19 +791,37 @@ pub fn decodeBlock(
870791 const block_size = block_header.block_size;
871792 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
872793 switch (block_header.block_type) {
873 .raw => return decodeRawBlock(dest[written_count..], src, block_size, consumed_count),
874 .rle => return decodeRleBlock(dest[written_count..], src, block_size, consumed_count),
794 .raw => {
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 },
875810 .compressed => {
876811 if (src.len < block_size) return error.MalformedBlockSize;
877812 var bytes_read: usize = 0;
878 const literals = decodeLiteralsSection(src, &bytes_read) catch
813 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
879814 return error.MalformedCompressedBlock;
880 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
815 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
816 const fbs_reader = fbs.reader();
817 const sequences_header = decodeSequencesHeader(fbs_reader) catch
881818 return error.MalformedCompressedBlock;
882819
883 bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch
820 decode_state.prepare(fbs_reader, literals, sequences_header) catch
884821 return error.MalformedCompressedBlock;
885822
823 bytes_read += fbs.pos;
824
886825 var bytes_written: usize = 0;
887826 if (sequences_header.sequence_count > 0) {
888827 const bit_stream_bytes = src[bytes_read..block_size];
......@@ -938,19 +877,37 @@ pub fn decodeBlockRingBuffer(
938877 const block_size = block_header.block_size;
939878 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
940879 switch (block_header.block_type) {
941 .raw => return decodeRawBlockRingBuffer(dest, src, block_size, consumed_count),
942 .rle => return decodeRleBlockRingBuffer(dest, src, block_size, consumed_count),
880 .raw => {
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 },
943896 .compressed => {
944897 if (src.len < block_size) return error.MalformedBlockSize;
945898 var bytes_read: usize = 0;
946 const literals = decodeLiteralsSection(src, &bytes_read) catch
899 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
947900 return error.MalformedCompressedBlock;
948 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
901 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
902 const fbs_reader = fbs.reader();
903 const sequences_header = decodeSequencesHeader(fbs_reader) catch
949904 return error.MalformedCompressedBlock;
950905
951 bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch
906 decode_state.prepare(fbs_reader, literals, sequences_header) catch
952907 return error.MalformedCompressedBlock;
953908
909 bytes_read += fbs.pos;
910
954911 var bytes_written: usize = 0;
955912 if (sequences_header.sequence_count > 0) {
956913 const bit_stream_bytes = src[bytes_read..block_size];
......@@ -991,6 +948,82 @@ pub fn decodeBlockRingBuffer(
991948 }
992949}
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
9941027/// Decode the header of a skippable frame.
9951028pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {
9961029 const magic = readInt(u32, src[0..4]);
......@@ -1002,13 +1035,6 @@ pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {
10021035 };
10031036}
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
10121038/// Returns the window size required to decompress a frame, or `null` if it cannot be
10131039/// determined, which indicates a malformed frame header.
10141040pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 {
......@@ -1023,40 +1049,37 @@ pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 {
10231049}
10241050
10251051const InvalidBit = error{ UnusedBitSet, ReservedBitSet };
1026/// Decode the header of a Zstandard frame. Returns `error.UnusedBitSet` or
1027/// `error.ReservedBitSet` if the corresponding bits are sets.
1028pub fn decodeZStandardHeader(src: []const u8, consumed_count: ?*usize) InvalidBit!frame.ZStandard.Header {
1029 const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, src[0]);
1052/// Decode the header of a Zstandard frame.
1053///
1054/// Errors:
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
10311061 if (descriptor.unused) return error.UnusedBitSet;
10321062 if (descriptor.reserved) return error.ReservedBitSet;
10331063
1034 var bytes_read_count: usize = 1;
1035
10361064 var window_descriptor: ?u8 = null;
10371065 if (!descriptor.single_segment_flag) {
1038 window_descriptor = src[bytes_read_count];
1039 bytes_read_count += 1;
1066 window_descriptor = try source.readByte();
10401067 }
10411068
10421069 var dictionary_id: ?u32 = null;
10431070 if (descriptor.dictionary_id_flag > 0) {
10441071 // if flag is 3 then field_size = 4, else field_size = flag
10451072 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]);
1047 bytes_read_count += field_size;
1073 dictionary_id = try source.readVarInt(u32, .Little, field_size);
10481074 }
10491075
10501076 var content_size: ?u64 = null;
10511077 if (descriptor.single_segment_flag or descriptor.content_size_flag > 0) {
10521078 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);
10541080 if (field_size == 2) content_size.? += 256;
1055 bytes_read_count += field_size;
10561081 }
10571082
1058 if (consumed_count) |p| p.* += bytes_read_count;
1059
10601083 const header = frame.ZStandard.Header{
10611084 .descriptor = descriptor,
10621085 .window_descriptor = window_descriptor,
......@@ -1080,12 +1103,20 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {
10801103
10811104/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
10821105/// 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(
10841111 src: []const u8,
10851112 consumed_count: *usize,
1086) (error{ MalformedLiteralsHeader, MalformedLiteralsSection } || DecodeHuffmanError)!LiteralsSection {
1113) (error{ MalformedLiteralsHeader, MalformedLiteralsSection, EndOfStream } || DecodeHuffmanError)!LiteralsSection {
10871114 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 };
10891120 switch (header.block_type) {
10901121 .raw => {
10911122 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;
......@@ -1110,7 +1141,7 @@ pub fn decodeLiteralsSection(
11101141 .compressed, .treeless => {
11111142 const huffman_tree_start = bytes_read;
11121143 const huffman_tree = if (header.block_type == .compressed)
1113 try decodeHuffmanTree(src[bytes_read..], &bytes_read)
1144 try decodeHuffmanTreeSlice(src[bytes_read..], &bytes_read)
11141145 else
11151146 null;
11161147 const huffman_tree_size = bytes_read - huffman_tree_start;
......@@ -1119,137 +1150,185 @@ pub fn decodeLiteralsSection(
11191150 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;
11201151 const stream_data = src[bytes_read .. bytes_read + total_streams_size];
11211152
1122 if (header.size_format == 0) {
1123 consumed_count.* += total_streams_size + bytes_read;
1124 return LiteralsSection{
1125 .header = header,
1126 .huffman_tree = huffman_tree,
1127 .streams = .{ .one = stream_data },
1128 };
1129 }
1130
1131 if (stream_data.len < 6) return error.MalformedLiteralsSection;
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);
1153 const streams = try decodeStreams(header.size_format, stream_data);
1154 consumed_count.* += bytes_read + total_streams_size;
1155 return LiteralsSection{
1156 .header = header,
1157 .huffman_tree = huffman_tree,
1158 .streams = streams,
1159 };
1160 },
1161 }
1162}
11371163
1138 const stream_1_start = 6;
1139 const stream_2_start = stream_1_start + stream_1_length;
1140 const stream_3_start = stream_2_start + stream_2_length;
1141 const stream_4_start = stream_3_start + stream_3_length;
1164/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
1165/// number of bytes the section uses.
1166///
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;
1144 consumed_count.* += total_streams_size + bytes_read;
1201 if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall;
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);
11461206 return LiteralsSection{
11471207 .header = header,
11481208 .huffman_tree = huffman_tree,
1149 .streams = .{ .four = .{
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 } },
1209 .streams = streams,
11551210 };
11561211 },
11571212 }
11581213}
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
11601239const DecodeHuffmanError = error{
11611240 MalformedHuffmanTree,
11621241 MalformedFseTable,
11631242 MalformedAccuracyLog,
11641243};
11651244
1166fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError!LiteralsSection.HuffmanTree {
1167 var bytes_read: usize = 0;
1168 bytes_read += 1;
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;
1245fn decodeFseHuffmanTree(source: anytype, compressed_size: usize, buffer: []u8, weights: *[256]u4) !usize {
1246 var stream = std.io.limitedReader(source, compressed_size);
1247 var bit_reader = bitReader(stream.reader());
12101248
1211 read_bits = 0;
1212 const odd_data = entries[odd_state];
1213 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;
1214 weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree;
1215 i += 1;
1216 if (read_bits < odd_data.bits) {
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;
1249 var entries: [1 << 6]Table.Fse = undefined;
1250 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
1251 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
1252 error.EndOfStream => return error.MalformedFseTable,
1253 };
1254 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
12241255
1225 symbol_count = i + 1; // stream contains all but the last symbol
1226 bytes_read += compressed_size;
1227 } else {
1228 const encoded_symbol_count = header - 127;
1229 symbol_count = encoded_symbol_count + 1;
1230 const weights_byte_count = (encoded_symbol_count + 1) / 2;
1231 if (src.len < weights_byte_count) return error.MalformedHuffmanTree;
1232 var i: usize = 0;
1233 while (i < weights_byte_count) : (i += 1) {
1234 weights[2 * i] = @intCast(u4, src[i + 1] >> 4);
1235 weights[2 * i + 1] = @intCast(u4, src[i + 1] & 0xF);
1256 const amount = try stream.reader().readAll(buffer);
1257 var huff_bits: ReverseBitReader = undefined;
1258 huff_bits.init(buffer[0..amount]) catch return error.MalformedHuffmanTree;
1259
1260 return assignWeights(&huff_bits, accuracy_log, &entries, weights);
1261}
1262
1263fn decodeFseHuffmanTreeSlice(src: []const u8, compressed_size: usize, weights: *[256]u4) !usize {
1264 if (src.len < compressed_size) return error.MalformedHuffmanTree;
1265 var stream = std.io.fixedBufferStream(src[0..compressed_size]);
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;
12361299 }
1237 bytes_read += weights_byte_count;
1238 }
1239 var weight_power_sum: u16 = 0;
1240 for (weights[0 .. symbol_count - 1]) |value| {
1241 if (value > 0) {
1242 weight_power_sum += @as(u16, 1) << (value - 1);
1300 even_state = even_data.baseline + even_bits;
1301
1302 read_bits = 0;
1303 const odd_data = entries[odd_state];
1304 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;
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;
12431312 }
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)
1247 max_number_of_bits = std.math.log2_int(u16, weight_power_sum) + 1;
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;
1316 return i + 1; // stream contains all but the last symbol
1317}
12501318
1251 var weight_sorted_prefixed_symbols: [256]LiteralsSection.HuffmanTree.PrefixedSymbol = undefined;
1252 for (weight_sorted_prefixed_symbols[0..symbol_count]) |_, i| {
1319fn decodeDirectHuffmanTree(source: anytype, encoded_symbol_count: usize, weights: *[256]u4) !usize {
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| {
12531332 weight_sorted_prefixed_symbols[i] = .{
12541333 .symbol = @intCast(u8, i),
12551334 .weight = undefined,
......@@ -1259,7 +1338,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
12591338
12601339 std.sort.sort(
12611340 LiteralsSection.HuffmanTree.PrefixedSymbol,
1262 weight_sorted_prefixed_symbols[0..symbol_count],
1341 weight_sorted_prefixed_symbols,
12631342 weights,
12641343 lessThanByWeight,
12651344 );
......@@ -1267,6 +1346,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
12671346 var prefix: u16 = 0;
12681347 var prefixed_symbol_count: usize = 0;
12691348 var sorted_index: usize = 0;
1349 const symbol_count = weight_sorted_prefixed_symbols.len;
12701350 while (sorted_index < symbol_count) {
12711351 var symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
12721352 const weight = weights[symbol];
......@@ -1290,7 +1370,24 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
12901370 weight_sorted_prefixed_symbols[prefixed_symbol_count].weight = weight;
12911371 }
12921372 }
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.*);
12941391 const tree = LiteralsSection.HuffmanTree{
12951392 .max_bit_count = max_number_of_bits,
12961393 .symbol_count_minus_one = @intCast(u8, prefixed_symbol_count - 1),
......@@ -1299,6 +1396,37 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError
12991396 return tree;
13001397}
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
13021430fn lessThanByWeight(
13031431 weights: [256]u4,
13041432 lhs: LiteralsSection.HuffmanTree.PrefixedSymbol,
......@@ -1311,9 +1439,8 @@ fn lessThanByWeight(
13111439}
13121440
13131441/// Decode a literals section header.
1314pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{MalformedLiteralsHeader}!LiteralsSection.Header {
1315 if (src.len == 0) return error.MalformedLiteralsHeader;
1316 const byte0 = src[0];
1442pub fn decodeLiteralsHeader(source: anytype) !LiteralsSection.Header {
1443 const byte0 = try source.readByte();
13171444 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);
13181445 const size_format = @intCast(u2, (byte0 & 0b1100) >> 2);
13191446 var regenerated_size: u20 = undefined;
......@@ -1323,47 +1450,31 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{Malfo
13231450 switch (size_format) {
13241451 0, 2 => {
13251452 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;
13401453 },
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),
13411458 }
13421459 },
13431460 .compressed, .treeless => {
1344 const byte1 = src[1];
1345 const byte2 = src[2];
1461 const byte1 = try source.readByte();
1462 const byte2 = try source.readByte();
13461463 switch (size_format) {
13471464 0, 1 => {
1348 if (src.len < 3) return error.MalformedLiteralsHeader;
13491465 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);
13501466 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);
1351 consumed_count.* += 3;
13521467 },
13531468 2 => {
1354 if (src.len < 4) return error.MalformedLiteralsHeader;
1355 const byte3 = src[3];
1469 const byte3 = try source.readByte();
13561470 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);
13571471 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);
1358 consumed_count.* += 4;
13591472 },
13601473 3 => {
1361 if (src.len < 5) return error.MalformedLiteralsHeader;
1362 const byte3 = src[3];
1363 const byte4 = src[4];
1474 const byte3 = try source.readByte();
1475 const byte4 = try source.readByte();
13641476 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);
13651477 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);
1366 consumed_count.* += 5;
13671478 },
13681479 }
13691480 },
......@@ -1377,18 +1488,17 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{Malfo
13771488}
13781489
13791490/// 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
13801495pub fn decodeSequencesHeader(
1381 src: []const u8,
1382 consumed_count: *usize,
1383) error{ MalformedSequencesHeader, ReservedBitSet }!SequencesSection.Header {
1384 if (src.len == 0) return error.MalformedSequencesHeader;
1496 source: anytype,
1497) !SequencesSection.Header {
13851498 var sequence_count: u24 = undefined;
13861499
1387 var bytes_read: usize = 0;
1388 const byte0 = src[0];
1500 const byte0 = try source.readByte();
13891501 if (byte0 == 0) {
1390 bytes_read += 1;
1391 consumed_count.* += bytes_read;
13921502 return SequencesSection.Header{
13931503 .sequence_count = 0,
13941504 .offsets = undefined,
......@@ -1397,22 +1507,14 @@ pub fn decodeSequencesHeader(
13971507 };
13981508 } else if (byte0 < 128) {
13991509 sequence_count = byte0;
1400 bytes_read += 1;
14011510 } else if (byte0 < 255) {
1402 if (src.len < 2) return error.MalformedSequencesHeader;
1403 sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1];
1404 bytes_read += 2;
1511 sequence_count = (@as(u24, (byte0 - 128)) << 8) + try source.readByte();
14051512 } else {
1406 if (src.len < 3) return error.MalformedSequencesHeader;
1407 sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00;
1408 bytes_read += 3;
1513 sequence_count = (try source.readByte()) + (@as(u24, try source.readByte()) << 8) + 0x7F00;
14091514 }
14101515
1411 if (src.len < bytes_read + 1) return error.MalformedSequencesHeader;
1412 const compression_modes = src[bytes_read];
1413 bytes_read += 1;
1516 const compression_modes = try source.readByte();
14141517
1415 consumed_count.* += bytes_read;
14161518 const matches_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00001100) >> 2);
14171519 const offsets_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00110000) >> 4);
14181520 const literal_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b11000000) >> 6);
......@@ -1615,7 +1717,7 @@ fn BitReader(comptime Reader: type) type {
16151717 };
16161718}
16171719
1618fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) {
1720pub fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) {
16191721 return .{ .underlying = std.io.bitReader(.Little, reader) };
16201722}
16211723