| ... | ... | @@ -14,29 +14,18 @@ fn readVarInt(comptime T: type, bytes: []const u8) T { |
| 14 | 14 | return std.mem.readVarInt(T, bytes, .Little); |
| 15 | 15 | } |
| 16 | 16 | |
| 17 | | fn isSkippableMagic(magic: u32) bool { |
| 17 | pub fn isSkippableMagic(magic: u32) bool { |
| 18 | 18 | return frame.Skippable.magic_number_min <= magic and magic <= frame.Skippable.magic_number_max; |
| 19 | 19 | } |
| 20 | 20 | |
| 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. |
| 25 | | pub 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. |
| 38 | | pub 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. |
| 27 | pub fn decodeFrameType(source: anytype) !frame.Kind { |
| 28 | const magic = try source.readIntLittle(u32); |
| 40 | 29 | return if (magic == frame.ZStandard.magic_number) |
| 41 | 30 | .zstandard |
| 42 | 31 | else if (isSkippableMagic(magic)) |
| ... | ... | @@ -52,15 +41,21 @@ const ReadWriteCount = struct { |
| 52 | 41 | |
| 53 | 42 | /// Decodes the frame at the start of `src` into `dest`. Returns the number of |
| 54 | 43 | /// bytes read from `src` and written to `dest`. |
| 44 | /// |
| 45 | /// Errors: |
| 46 | /// - returns `error.UnknownContentSizeUnsupported` |
| 47 | /// - returns `error.ContentTooLarge` |
| 48 | /// - returns `error.BadMagic` |
| 55 | 49 | pub fn decodeFrame( |
| 56 | 50 | dest: []u8, |
| 57 | 51 | src: []const u8, |
| 58 | 52 | 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())) { |
| 61 | 56 | .zstandard => decodeZStandardFrame(dest, src, verify_checksum), |
| 62 | 57 | .skippable => ReadWriteCount{ |
| 63 | | .read_count = skippableFrameSize(src[0..8]) + 8, |
| 58 | .read_count = try fbs.reader().readIntLittle(u32) + 8, |
| 64 | 59 | .write_count = 0, |
| 65 | 60 | }, |
| 66 | 61 | }; |
| ... | ... | @@ -97,16 +92,52 @@ pub const DecodeState = struct { |
| 97 | 92 | }; |
| 98 | 93 | } |
| 99 | 94 | |
| 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 | 126 | /// 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. |
| 104 | 135 | pub fn prepare( |
| 105 | 136 | self: *DecodeState, |
| 106 | | src: []const u8, |
| 137 | source: anytype, |
| 107 | 138 | literals: LiteralsSection, |
| 108 | 139 | sequences_header: SequencesSection.Header, |
| 109 | | ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize { |
| 140 | ) !void { |
| 110 | 141 | self.literal_written_count = 0; |
| 111 | 142 | self.literal_header = literals.header; |
| 112 | 143 | self.literal_streams = literals.streams; |
| ... | ... | @@ -129,28 +160,11 @@ pub const DecodeState = struct { |
| 129 | 160 | } |
| 130 | 161 | |
| 131 | 162 | 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); |
| 149 | 166 | self.fse_tables_undefined = false; |
| 150 | | |
| 151 | | return bytes_read; |
| 152 | 167 | } |
| 153 | | return 0; |
| 154 | 168 | } |
| 155 | 169 | |
| 156 | 170 | /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream` |
| ... | ... | @@ -208,10 +222,10 @@ pub const DecodeState = struct { |
| 208 | 222 | |
| 209 | 223 | fn updateFseTable( |
| 210 | 224 | self: *DecodeState, |
| 211 | | src: []const u8, |
| 225 | source: anytype, |
| 212 | 226 | comptime choice: DataType, |
| 213 | 227 | mode: SequencesSection.Header.Mode, |
| 214 | | ) FseTableError!usize { |
| 228 | ) !void { |
| 215 | 229 | const field_name = @tagName(choice); |
| 216 | 230 | switch (mode) { |
| 217 | 231 | .predefined => { |
| ... | ... | @@ -220,17 +234,13 @@ pub const DecodeState = struct { |
| 220 | 234 | |
| 221 | 235 | @field(self, field_name).table = |
| 222 | 236 | @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table"); |
| 223 | | return 0; |
| 224 | 237 | }, |
| 225 | 238 | .rle => { |
| 226 | 239 | @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() }; |
| 229 | 241 | }, |
| 230 | 242 | .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); |
| 234 | 244 | |
| 235 | 245 | const table_size = try decodeFseTable( |
| 236 | 246 | &bit_reader, |
| ... | ... | @@ -242,9 +252,8 @@ pub const DecodeState = struct { |
| 242 | 252 | .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size], |
| 243 | 253 | }; |
| 244 | 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 | } |
| 250 | 259 | |
| ... | ... | @@ -462,11 +471,15 @@ pub const DecodeState = struct { |
| 462 | 471 | while (i < len) : (i += 1) { |
| 463 | 472 | var prefix: u16 = 0; |
| 464 | 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 | 477 | prefix <<= bit_count_to_read; |
| 467 | 478 | prefix |= new_bits; |
| 468 | 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 | }; |
| 470 | 483 | |
| 471 | 484 | switch (result) { |
| 472 | 485 | .symbol => |sym| { |
| ... | ... | @@ -589,11 +602,14 @@ pub fn decodeZStandardFrame( |
| 589 | 602 | dest: []u8, |
| 590 | 603 | src: []const u8, |
| 591 | 604 | verify_checksum: bool, |
| 592 | | ) (error{ UnknownContentSizeUnsupported, ContentTooLarge } || FrameError)!ReadWriteCount { |
| 605 | ) (error{ UnknownContentSizeUnsupported, ContentTooLarge, EndOfStream } || FrameError)!ReadWriteCount { |
| 593 | 606 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); |
| 594 | 607 | var consumed_count: usize = 4; |
| 595 | 608 | |
| 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; |
| 597 | 613 | |
| 598 | 614 | if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported; |
| 599 | 615 | |
| ... | ... | @@ -649,18 +665,25 @@ pub const FrameContext = struct { |
| 649 | 665 | /// `decodeZStandardFrame()`. Returns `error.WindowSizeUnknown` if the frame |
| 650 | 666 | /// does not declare its content size or a window descriptor (this indicates a |
| 651 | 667 | /// malformed frame). |
| 668 | /// |
| 669 | /// Errors: |
| 670 | /// - returns `error.WindowTooLarge` |
| 671 | /// - returns `error.WindowSizeUnknown` |
| 652 | 672 | pub fn decodeZStandardFrameAlloc( |
| 653 | 673 | allocator: std.mem.Allocator, |
| 654 | 674 | src: []const u8, |
| 655 | 675 | verify_checksum: bool, |
| 656 | 676 | window_size_max: usize, |
| 657 | | ) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory } || FrameError)![]u8 { |
| 677 | ) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory, EndOfStream } || FrameError)![]u8 { |
| 658 | 678 | var result = std.ArrayList(u8).init(allocator); |
| 659 | 679 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); |
| 660 | 680 | var consumed_count: usize = 4; |
| 661 | 681 | |
| 662 | 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 | 687 | break :context try FrameContext.init(frame_header, window_size_max, verify_checksum); |
| 665 | 688 | }; |
| 666 | 689 | |
| ... | ... | @@ -674,30 +697,7 @@ pub fn decodeZStandardFrameAlloc( |
| 674 | 697 | |
| 675 | 698 | var block_header = decodeBlockHeader(src[consumed_count..][0..3]); |
| 676 | 699 | 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); |
| 701 | 701 | while (true) : ({ |
| 702 | 702 | block_header = decodeBlockHeader(src[consumed_count..][0..3]); |
| 703 | 703 | consumed_count += 3; |
| ... | ... | @@ -754,30 +754,7 @@ pub fn decodeFrameBlocks( |
| 754 | 754 | var block_header = decodeBlockHeader(src[0..3]); |
| 755 | 755 | var bytes_read: usize = 3; |
| 756 | 756 | 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); |
| 781 | 758 | var written_count: usize = 0; |
| 782 | 759 | while (true) : ({ |
| 783 | 760 | block_header = decodeBlockHeader(src[bytes_read..][0..3]); |
| ... | ... | @@ -798,62 +775,6 @@ pub fn decodeFrameBlocks( |
| 798 | 775 | return written_count; |
| 799 | 776 | } |
| 800 | 777 | |
| 801 | | fn 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 | | |
| 814 | | fn 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 | | |
| 827 | | fn 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 | | |
| 842 | | fn 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 | 778 | /// Decode a single block from `src` into `dest`. The beginning of `src` should |
| 858 | 779 | /// be the start of the block content (i.e. directly after the block header). |
| 859 | 780 | /// Increments `consumed_count` by the number of bytes read from `src` to decode |
| ... | ... | @@ -870,19 +791,37 @@ pub fn decodeBlock( |
| 870 | 791 | const block_size = block_header.block_size; |
| 871 | 792 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; |
| 872 | 793 | 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 | }, |
| 875 | 810 | .compressed => { |
| 876 | 811 | if (src.len < block_size) return error.MalformedBlockSize; |
| 877 | 812 | var bytes_read: usize = 0; |
| 878 | | const literals = decodeLiteralsSection(src, &bytes_read) catch |
| 813 | const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch |
| 879 | 814 | 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 |
| 881 | 818 | return error.MalformedCompressedBlock; |
| 882 | 819 | |
| 883 | | bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch |
| 820 | decode_state.prepare(fbs_reader, literals, sequences_header) catch |
| 884 | 821 | return error.MalformedCompressedBlock; |
| 885 | 822 | |
| 823 | bytes_read += fbs.pos; |
| 824 | |
| 886 | 825 | var bytes_written: usize = 0; |
| 887 | 826 | if (sequences_header.sequence_count > 0) { |
| 888 | 827 | const bit_stream_bytes = src[bytes_read..block_size]; |
| ... | ... | @@ -938,19 +877,37 @@ pub fn decodeBlockRingBuffer( |
| 938 | 877 | const block_size = block_header.block_size; |
| 939 | 878 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; |
| 940 | 879 | 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 | }, |
| 943 | 896 | .compressed => { |
| 944 | 897 | if (src.len < block_size) return error.MalformedBlockSize; |
| 945 | 898 | var bytes_read: usize = 0; |
| 946 | | const literals = decodeLiteralsSection(src, &bytes_read) catch |
| 899 | const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch |
| 947 | 900 | 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 |
| 949 | 904 | return error.MalformedCompressedBlock; |
| 950 | 905 | |
| 951 | | bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch |
| 906 | decode_state.prepare(fbs_reader, literals, sequences_header) catch |
| 952 | 907 | return error.MalformedCompressedBlock; |
| 953 | 908 | |
| 909 | bytes_read += fbs.pos; |
| 910 | |
| 954 | 911 | var bytes_written: usize = 0; |
| 955 | 912 | if (sequences_header.sequence_count > 0) { |
| 956 | 913 | const bit_stream_bytes = src[bytes_read..block_size]; |
| ... | ... | @@ -991,6 +948,82 @@ pub fn decodeBlockRingBuffer( |
| 991 | 948 | } |
| 992 | 949 | } |
| 993 | 950 | |
| 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. |
| 958 | pub 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 | 1027 | /// Decode the header of a skippable frame. |
| 995 | 1028 | pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header { |
| 996 | 1029 | const magic = readInt(u32, src[0..4]); |
| ... | ... | @@ -1002,13 +1035,6 @@ pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header { |
| 1002 | 1035 | }; |
| 1003 | 1036 | } |
| 1004 | 1037 | |
| 1005 | | /// Returns the content size of a skippable frame. |
| 1006 | | pub 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 | 1038 | /// Returns the window size required to decompress a frame, or `null` if it cannot be |
| 1013 | 1039 | /// determined, which indicates a malformed frame header. |
| 1014 | 1040 | pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 { |
| ... | ... | @@ -1023,40 +1049,37 @@ pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 { |
| 1023 | 1049 | } |
| 1024 | 1050 | |
| 1025 | 1051 | const 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. |
| 1028 | | pub 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 |
| 1058 | pub fn decodeZStandardHeader(source: anytype) (error{EndOfStream} || InvalidBit)!frame.ZStandard.Header { |
| 1059 | const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, try source.readByte()); |
| 1030 | 1060 | |
| 1031 | 1061 | if (descriptor.unused) return error.UnusedBitSet; |
| 1032 | 1062 | if (descriptor.reserved) return error.ReservedBitSet; |
| 1033 | 1063 | |
| 1034 | | var bytes_read_count: usize = 1; |
| 1035 | | |
| 1036 | 1064 | var window_descriptor: ?u8 = null; |
| 1037 | 1065 | if (!descriptor.single_segment_flag) { |
| 1038 | | window_descriptor = src[bytes_read_count]; |
| 1039 | | bytes_read_count += 1; |
| 1066 | window_descriptor = try source.readByte(); |
| 1040 | 1067 | } |
| 1041 | 1068 | |
| 1042 | 1069 | var dictionary_id: ?u32 = null; |
| 1043 | 1070 | if (descriptor.dictionary_id_flag > 0) { |
| 1044 | 1071 | // if flag is 3 then field_size = 4, else field_size = flag |
| 1045 | 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]); |
| 1047 | | bytes_read_count += field_size; |
| 1073 | dictionary_id = try source.readVarInt(u32, .Little, field_size); |
| 1048 | 1074 | } |
| 1049 | 1075 | |
| 1050 | 1076 | var content_size: ?u64 = null; |
| 1051 | 1077 | if (descriptor.single_segment_flag or descriptor.content_size_flag > 0) { |
| 1052 | 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 | 1080 | if (field_size == 2) content_size.? += 256; |
| 1055 | | bytes_read_count += field_size; |
| 1056 | 1081 | } |
| 1057 | 1082 | |
| 1058 | | if (consumed_count) |p| p.* += bytes_read_count; |
| 1059 | | |
| 1060 | 1083 | const header = frame.ZStandard.Header{ |
| 1061 | 1084 | .descriptor = descriptor, |
| 1062 | 1085 | .window_descriptor = window_descriptor, |
| ... | ... | @@ -1080,12 +1103,20 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header { |
| 1080 | 1103 | |
| 1081 | 1104 | /// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the |
| 1082 | 1105 | /// number of bytes the section uses. |
| 1083 | | pub fn decodeLiteralsSection( |
| 1106 | /// |
| 1107 | /// Errors: |
| 1108 | /// - returns `error.MalformedLiteralsHeader` if the header is invalid |
| 1109 | /// - returns `error.MalformedLiteralsSection` if there are errors decoding |
| 1110 | pub fn decodeLiteralsSectionSlice( |
| 1084 | 1111 | src: []const u8, |
| 1085 | 1112 | consumed_count: *usize, |
| 1086 | | ) (error{ MalformedLiteralsHeader, MalformedLiteralsSection } || DecodeHuffmanError)!LiteralsSection { |
| 1113 | ) (error{ MalformedLiteralsHeader, MalformedLiteralsSection, EndOfStream } || DecodeHuffmanError)!LiteralsSection { |
| 1087 | 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 | 1120 | switch (header.block_type) { |
| 1090 | 1121 | .raw => { |
| 1091 | 1122 | if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection; |
| ... | ... | @@ -1110,7 +1141,7 @@ pub fn decodeLiteralsSection( |
| 1110 | 1141 | .compressed, .treeless => { |
| 1111 | 1142 | const huffman_tree_start = bytes_read; |
| 1112 | 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 | 1145 | else |
| 1115 | 1146 | null; |
| 1116 | 1147 | const huffman_tree_size = bytes_read - huffman_tree_start; |
| ... | ... | @@ -1119,137 +1150,185 @@ pub fn decodeLiteralsSection( |
| 1119 | 1150 | if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection; |
| 1120 | 1151 | const stream_data = src[bytes_read .. bytes_read + total_streams_size]; |
| 1121 | 1152 | |
| 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 | } |
| 1137 | 1163 | |
| 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 |
| 1170 | pub 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); |
| 1142 | 1200 | |
| 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]; |
| 1145 | 1204 | |
| 1205 | const streams = try decodeStreams(header.size_format, stream_data); |
| 1146 | 1206 | return LiteralsSection{ |
| 1147 | 1207 | .header = header, |
| 1148 | 1208 | .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, |
| 1155 | 1210 | }; |
| 1156 | 1211 | }, |
| 1157 | 1212 | } |
| 1158 | 1213 | } |
| 1159 | 1214 | |
| 1215 | fn 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 | |
| 1160 | 1239 | const DecodeHuffmanError = error{ |
| 1161 | 1240 | MalformedHuffmanTree, |
| 1162 | 1241 | MalformedFseTable, |
| 1163 | 1242 | MalformedAccuracyLog, |
| 1164 | 1243 | }; |
| 1165 | 1244 | |
| 1166 | | fn 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; |
| 1245 | fn 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()); |
| 1210 | 1248 | |
| 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); |
| 1224 | 1255 | |
| 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 | |
| 1263 | fn 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 | |
| 1284 | fn 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; |
| 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; |
| 1243 | 1312 | } |
| 1244 | | } |
| 1313 | odd_state = odd_data.baseline + odd_bits; |
| 1314 | } else return error.MalformedHuffmanTree; |
| 1245 | 1315 | |
| 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 | } |
| 1250 | 1318 | |
| 1251 | | var weight_sorted_prefixed_symbols: [256]LiteralsSection.HuffmanTree.PrefixedSymbol = undefined; |
| 1252 | | for (weight_sorted_prefixed_symbols[0..symbol_count]) |_, i| { |
| 1319 | fn 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 | |
| 1330 | fn assignSymbols(weight_sorted_prefixed_symbols: []LiteralsSection.HuffmanTree.PrefixedSymbol, weights: [256]u4) usize { |
| 1331 | for (weight_sorted_prefixed_symbols) |_, i| { |
| 1253 | 1332 | weight_sorted_prefixed_symbols[i] = .{ |
| 1254 | 1333 | .symbol = @intCast(u8, i), |
| 1255 | 1334 | .weight = undefined, |
| ... | ... | @@ -1259,7 +1338,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError |
| 1259 | 1338 | |
| 1260 | 1339 | std.sort.sort( |
| 1261 | 1340 | LiteralsSection.HuffmanTree.PrefixedSymbol, |
| 1262 | | weight_sorted_prefixed_symbols[0..symbol_count], |
| 1341 | weight_sorted_prefixed_symbols, |
| 1263 | 1342 | weights, |
| 1264 | 1343 | lessThanByWeight, |
| 1265 | 1344 | ); |
| ... | ... | @@ -1267,6 +1346,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError |
| 1267 | 1346 | var prefix: u16 = 0; |
| 1268 | 1347 | var prefixed_symbol_count: usize = 0; |
| 1269 | 1348 | var sorted_index: usize = 0; |
| 1349 | const symbol_count = weight_sorted_prefixed_symbols.len; |
| 1270 | 1350 | while (sorted_index < symbol_count) { |
| 1271 | 1351 | var symbol = weight_sorted_prefixed_symbols[sorted_index].symbol; |
| 1272 | 1352 | const weight = weights[symbol]; |
| ... | ... | @@ -1290,7 +1370,24 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError |
| 1290 | 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 | |
| 1376 | fn 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 | 1391 | const tree = LiteralsSection.HuffmanTree{ |
| 1295 | 1392 | .max_bit_count = max_number_of_bits, |
| 1296 | 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 | 1396 | return tree; |
| 1300 | 1397 | } |
| 1301 | 1398 | |
| 1399 | fn 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 | |
| 1411 | fn 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 | |
| 1302 | 1430 | fn lessThanByWeight( |
| 1303 | 1431 | weights: [256]u4, |
| 1304 | 1432 | lhs: LiteralsSection.HuffmanTree.PrefixedSymbol, |
| ... | ... | @@ -1311,9 +1439,8 @@ fn lessThanByWeight( |
| 1311 | 1439 | } |
| 1312 | 1440 | |
| 1313 | 1441 | /// Decode a literals section header. |
| 1314 | | pub 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]; |
| 1442 | pub fn decodeLiteralsHeader(source: anytype) !LiteralsSection.Header { |
| 1443 | const byte0 = try source.readByte(); |
| 1317 | 1444 | const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11); |
| 1318 | 1445 | const size_format = @intCast(u2, (byte0 & 0b1100) >> 2); |
| 1319 | 1446 | var regenerated_size: u20 = undefined; |
| ... | ... | @@ -1323,47 +1450,31 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{Malfo |
| 1323 | 1450 | switch (size_format) { |
| 1324 | 1451 | 0, 2 => { |
| 1325 | 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 | 1460 | .compressed, .treeless => { |
| 1344 | | const byte1 = src[1]; |
| 1345 | | const byte2 = src[2]; |
| 1461 | const byte1 = try source.readByte(); |
| 1462 | const byte2 = try source.readByte(); |
| 1346 | 1463 | switch (size_format) { |
| 1347 | 1464 | 0, 1 => { |
| 1348 | | if (src.len < 3) return error.MalformedLiteralsHeader; |
| 1349 | 1465 | regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4); |
| 1350 | 1466 | compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2); |
| 1351 | | consumed_count.* += 3; |
| 1352 | 1467 | }, |
| 1353 | 1468 | 2 => { |
| 1354 | | if (src.len < 4) return error.MalformedLiteralsHeader; |
| 1355 | | const byte3 = src[3]; |
| 1469 | const byte3 = try source.readByte(); |
| 1356 | 1470 | regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12); |
| 1357 | 1471 | compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6); |
| 1358 | | consumed_count.* += 4; |
| 1359 | 1472 | }, |
| 1360 | 1473 | 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(); |
| 1364 | 1476 | regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12); |
| 1365 | 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 | 1488 | } |
| 1378 | 1489 | |
| 1379 | 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 |
| 1380 | 1495 | pub 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 { |
| 1385 | 1498 | var sequence_count: u24 = undefined; |
| 1386 | 1499 | |
| 1387 | | var bytes_read: usize = 0; |
| 1388 | | const byte0 = src[0]; |
| 1500 | const byte0 = try source.readByte(); |
| 1389 | 1501 | if (byte0 == 0) { |
| 1390 | | bytes_read += 1; |
| 1391 | | consumed_count.* += bytes_read; |
| 1392 | 1502 | return SequencesSection.Header{ |
| 1393 | 1503 | .sequence_count = 0, |
| 1394 | 1504 | .offsets = undefined, |
| ... | ... | @@ -1397,22 +1507,14 @@ pub fn decodeSequencesHeader( |
| 1397 | 1507 | }; |
| 1398 | 1508 | } else if (byte0 < 128) { |
| 1399 | 1509 | sequence_count = byte0; |
| 1400 | | bytes_read += 1; |
| 1401 | 1510 | } 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(); |
| 1405 | 1512 | } 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; |
| 1409 | 1514 | } |
| 1410 | 1515 | |
| 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(); |
| 1414 | 1517 | |
| 1415 | | consumed_count.* += bytes_read; |
| 1416 | 1518 | const matches_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00001100) >> 2); |
| 1417 | 1519 | const offsets_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00110000) >> 4); |
| 1418 | 1520 | const literal_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b11000000) >> 6); |
| ... | ... | @@ -1615,7 +1717,7 @@ fn BitReader(comptime Reader: type) type { |
| 1615 | 1717 | }; |
| 1616 | 1718 | } |
| 1617 | 1719 | |
| 1618 | | fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) { |
| 1720 | pub fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) { |
| 1619 | 1721 | return .{ .underlying = std.io.bitReader(.Little, reader) }; |
| 1620 | 1722 | } |
| 1621 | 1723 | |