authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-01-31 13:24:27+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log947ad3e26816270fa11298dd5f118e50126a1320
treea608920838ef177ab8bb6d84f98fa986e5f6c6a3
parent2d35c16ee7e4a8f69bbbe19b2a48cc03aed755c8

std.compress.zstandard: add FrameContext and add literals into DecodeState


2 files changed, 83 insertions(+), 71 deletions(-)

lib/std/compress/zstandard/decompress.zig+78-71
......@@ -81,6 +81,8 @@ pub const DecodeState = struct {
8181
8282 literal_stream_reader: ReverseBitReader,
8383 literal_stream_index: usize,
84 literal_streams: LiteralsSection.Streams,
85 literal_header: LiteralsSection.Header,
8486 huffman_tree: ?LiteralsSection.HuffmanTree,
8587
8688 literal_written_count: usize,
......@@ -105,6 +107,10 @@ pub const DecodeState = struct {
105107 literals: LiteralsSection,
106108 sequences_header: SequencesSection.Header,
107109 ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize {
110 self.literal_written_count = 0;
111 self.literal_header = literals.header;
112 self.literal_streams = literals.streams;
113
108114 if (literals.huffman_tree) |tree| {
109115 self.huffman_tree = tree;
110116 } else if (literals.header.block_type == .treeless and self.huffman_tree == null) {
......@@ -293,12 +299,11 @@ pub const DecodeState = struct {
293299 self: *DecodeState,
294300 dest: []u8,
295301 write_pos: usize,
296 literals: LiteralsSection,
297302 sequence: Sequence,
298303 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
299304 if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence;
300305
301 try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length);
306 try self.decodeLiteralsSlice(dest[write_pos..], sequence.literal_length);
302307 const copy_start = write_pos + sequence.literal_length - sequence.offset;
303308 const copy_end = copy_start + sequence.match_length;
304309 // NOTE: we ignore the usage message for std.mem.copy and copy with dest.ptr >= src.ptr
......@@ -309,12 +314,11 @@ pub const DecodeState = struct {
309314 fn executeSequenceRingBuffer(
310315 self: *DecodeState,
311316 dest: *RingBuffer,
312 literals: LiteralsSection,
313317 sequence: Sequence,
314318 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
315319 if (sequence.offset > dest.data.len) return error.MalformedSequence;
316320
317 try self.decodeLiteralsRingBuffer(dest, literals, sequence.literal_length);
321 try self.decodeLiteralsRingBuffer(dest, sequence.literal_length);
318322 const copy_start = dest.write_index + dest.data.len - sequence.offset;
319323 const copy_slice = dest.sliceAt(copy_start, sequence.match_length);
320324 // TODO: would std.mem.copy and figuring out dest slice be better/faster?
......@@ -328,6 +332,7 @@ pub const DecodeState = struct {
328332 MalformedSequence,
329333 MalformedFseBits,
330334 } || DecodeLiteralsError;
335
331336 /// Decode one sequence from `bit_reader` into `dest`, written starting at
332337 /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns
333338 /// `error.MalformedSequence` error if the decompressed sequence would be longer
......@@ -340,7 +345,6 @@ pub const DecodeState = struct {
340345 self: *DecodeState,
341346 dest: []u8,
342347 write_pos: usize,
343 literals: LiteralsSection,
344348 bit_reader: *ReverseBitReader,
345349 sequence_size_limit: usize,
346350 last_sequence: bool,
......@@ -349,7 +353,7 @@ pub const DecodeState = struct {
349353 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
350354 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
351355
352 try self.executeSequenceSlice(dest, write_pos, literals, sequence);
356 try self.executeSequenceSlice(dest, write_pos, sequence);
353357 if (!last_sequence) {
354358 try self.updateState(.literal, bit_reader);
355359 try self.updateState(.match, bit_reader);
......@@ -362,7 +366,6 @@ pub const DecodeState = struct {
362366 pub fn decodeSequenceRingBuffer(
363367 self: *DecodeState,
364368 dest: *RingBuffer,
365 literals: LiteralsSection,
366369 bit_reader: anytype,
367370 sequence_size_limit: usize,
368371 last_sequence: bool,
......@@ -371,7 +374,7 @@ pub const DecodeState = struct {
371374 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
372375 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
373376
374 try self.executeSequenceRingBuffer(dest, literals, sequence);
377 try self.executeSequenceRingBuffer(dest, sequence);
375378 if (!last_sequence) {
376379 try self.updateState(.literal, bit_reader);
377380 try self.updateState(.match, bit_reader);
......@@ -382,13 +385,12 @@ pub const DecodeState = struct {
382385
383386 fn nextLiteralMultiStream(
384387 self: *DecodeState,
385 literals: LiteralsSection,
386388 ) error{BitStreamHasNoStartBit}!void {
387389 self.literal_stream_index += 1;
388 try self.initLiteralStream(literals.streams.four[self.literal_stream_index]);
390 try self.initLiteralStream(self.literal_streams.four[self.literal_stream_index]);
389391 }
390392
391 fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
393 pub fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
392394 try self.literal_stream_reader.init(bytes);
393395 }
394396
......@@ -400,11 +402,10 @@ pub const DecodeState = struct {
400402 self: *DecodeState,
401403 comptime T: type,
402404 bit_count_to_read: usize,
403 literals: LiteralsSection,
404405 ) LiteralBitsError!T {
405406 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {
406 if (literals.streams == .four and self.literal_stream_index < 3) {
407 try self.nextLiteralMultiStream(literals);
407 if (self.literal_streams == .four and self.literal_stream_index < 3) {
408 try self.nextLiteralMultiStream();
408409 break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch
409410 return error.UnexpectedEndOfLiteralStream;
410411 } else {
......@@ -427,23 +428,22 @@ pub const DecodeState = struct {
427428 pub fn decodeLiteralsSlice(
428429 self: *DecodeState,
429430 dest: []u8,
430 literals: LiteralsSection,
431431 len: usize,
432432 ) DecodeLiteralsError!void {
433 if (self.literal_written_count + len > literals.header.regenerated_size)
433 if (self.literal_written_count + len > self.literal_header.regenerated_size)
434434 return error.MalformedLiteralsLength;
435435
436 switch (literals.header.block_type) {
436 switch (self.literal_header.block_type) {
437437 .raw => {
438438 const literals_end = self.literal_written_count + len;
439 const literal_data = literals.streams.one[self.literal_written_count..literals_end];
439 const literal_data = self.literal_streams.one[self.literal_written_count..literals_end];
440440 std.mem.copy(u8, dest, literal_data);
441441 self.literal_written_count += len;
442442 },
443443 .rle => {
444444 var i: usize = 0;
445445 while (i < len) : (i += 1) {
446 dest[i] = literals.streams.one[0];
446 dest[i] = self.literal_streams.one[0];
447447 }
448448 self.literal_written_count += len;
449449 },
......@@ -462,7 +462,7 @@ pub const DecodeState = struct {
462462 while (i < len) : (i += 1) {
463463 var prefix: u16 = 0;
464464 while (true) {
465 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals);
465 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read);
466466 prefix <<= bit_count_to_read;
467467 prefix |= new_bits;
468468 bits_read += bit_count_to_read;
......@@ -496,23 +496,22 @@ pub const DecodeState = struct {
496496 pub fn decodeLiteralsRingBuffer(
497497 self: *DecodeState,
498498 dest: *RingBuffer,
499 literals: LiteralsSection,
500499 len: usize,
501500 ) DecodeLiteralsError!void {
502 if (self.literal_written_count + len > literals.header.regenerated_size)
501 if (self.literal_written_count + len > self.literal_header.regenerated_size)
503502 return error.MalformedLiteralsLength;
504503
505 switch (literals.header.block_type) {
504 switch (self.literal_header.block_type) {
506505 .raw => {
507506 const literals_end = self.literal_written_count + len;
508 const literal_data = literals.streams.one[self.literal_written_count..literals_end];
507 const literal_data = self.literal_streams.one[self.literal_written_count..literals_end];
509508 dest.writeSliceAssumeCapacity(literal_data);
510509 self.literal_written_count += len;
511510 },
512511 .rle => {
513512 var i: usize = 0;
514513 while (i < len) : (i += 1) {
515 dest.writeAssumeCapacity(literals.streams.one[0]);
514 dest.writeAssumeCapacity(self.literal_streams.one[0]);
516515 }
517516 self.literal_written_count += len;
518517 },
......@@ -531,7 +530,7 @@ pub const DecodeState = struct {
531530 while (i < len) : (i += 1) {
532531 var prefix: u16 = 0;
533532 while (true) {
534 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals);
533 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read);
535534 prefix <<= bit_count_to_read;
536535 prefix |= new_bits;
537536 bits_read += bit_count_to_read;
......@@ -569,10 +568,6 @@ pub const DecodeState = struct {
569568 }
570569};
571570
572const literal_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.literal;
573const match_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match;
574const offset_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match;
575
576571pub fn computeChecksum(hasher: *std.hash.XxHash64) u32 {
577572 const hash = hasher.final();
578573 return @intCast(u32, hash & 0xFFFFFFFF);
......@@ -625,6 +620,31 @@ pub fn decodeZStandardFrame(
625620 return ReadWriteCount{ .read_count = consumed_count, .write_count = written_count };
626621}
627622
623pub const FrameContext = struct {
624 hasher_opt: ?std.hash.XxHash64,
625 window_size: usize,
626 has_checksum: bool,
627 block_size_max: usize,
628
629 pub fn init(frame_header: frame.ZStandard.Header, window_size_max: usize, verify_checksum: bool) !FrameContext {
630 if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported;
631
632 const window_size_raw = frameWindowSize(frame_header) orelse return error.WindowSizeUnknown;
633 const window_size = if (window_size_raw > window_size_max)
634 return error.WindowTooLarge
635 else
636 @intCast(usize, window_size_raw);
637
638 const should_compute_checksum = frame_header.descriptor.content_checksum_flag and verify_checksum;
639 return .{
640 .hasher_opt = if (should_compute_checksum) std.hash.XxHash64.init(0) else null,
641 .window_size = window_size,
642 .has_checksum = frame_header.descriptor.content_checksum_flag,
643 .block_size_max = @min(1 << 17, window_size),
644 };
645 }
646};
647
628648/// Decode a Zstandard from from `src` and return the decompressed bytes; see
629649/// `decodeZStandardFrame()`. Returns `error.WindowSizeUnknown` if the frame
630650/// does not declare its content size or a window descriptor (this indicates a
......@@ -639,33 +659,18 @@ pub fn decodeZStandardFrameAlloc(
639659 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
640660 var consumed_count: usize = 4;
641661
642 const frame_header = try decodeZStandardHeader(src[consumed_count..], &consumed_count);
643
644 if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported;
645
646 const window_size_raw = frameWindowSize(frame_header) orelse return error.WindowSizeUnknown;
647 const window_size = if (window_size_raw > window_size_max)
648 return error.WindowTooLarge
649 else
650 @intCast(usize, window_size_raw);
651
652 const should_compute_checksum = frame_header.descriptor.content_checksum_flag and verify_checksum;
653 var hasher_opt = if (should_compute_checksum) std.hash.XxHash64.init(0) else null;
654
655 const block_size_maximum = @min(1 << 17, window_size);
656
657 var window_data = try allocator.alloc(u8, window_size);
658 defer allocator.free(window_data);
659 var ring_buffer = RingBuffer{
660 .data = window_data,
661 .write_index = 0,
662 .read_index = 0,
662 var frame_context = context: {
663 const frame_header = try decodeZStandardHeader(src[consumed_count..], &consumed_count);
664 break :context try FrameContext.init(frame_header, window_size_max, verify_checksum);
663665 };
664666
667 var ring_buffer = try RingBuffer.init(allocator, frame_context.window_size);
668 defer ring_buffer.deinit(allocator);
669
665670 // These tables take 7680 bytes
666 var literal_fse_data: [literal_table_size_max]Table.Fse = undefined;
667 var match_fse_data: [match_table_size_max]Table.Fse = undefined;
668 var offset_fse_data: [offset_table_size_max]Table.Fse = undefined;
671 var literal_fse_data: [types.compressed_block.table_size_max.literal]Table.Fse = undefined;
672 var match_fse_data: [types.compressed_block.table_size_max.match]Table.Fse = undefined;
673 var offset_fse_data: [types.compressed_block.table_size_max.offset]Table.Fse = undefined;
669674
670675 var block_header = decodeBlockHeader(src[consumed_count..][0..3]);
671676 consumed_count += 3;
......@@ -687,6 +692,8 @@ pub fn decodeZStandardFrameAlloc(
687692 .fse_tables_undefined = true,
688693
689694 .literal_written_count = 0,
695 .literal_header = undefined,
696 .literal_streams = undefined,
690697 .literal_stream_reader = undefined,
691698 .literal_stream_index = undefined,
692699 .huffman_tree = null,
......@@ -695,30 +702,29 @@ pub fn decodeZStandardFrameAlloc(
695702 block_header = decodeBlockHeader(src[consumed_count..][0..3]);
696703 consumed_count += 3;
697704 }) {
698 if (block_header.block_size > block_size_maximum) return error.BlockSizeOverMaximum;
705 if (block_header.block_size > frame_context.block_size_max) return error.BlockSizeOverMaximum;
699706 const written_size = try decodeBlockRingBuffer(
700707 &ring_buffer,
701708 src[consumed_count..],
702709 block_header,
703710 &decode_state,
704711 &consumed_count,
705 block_size_maximum,
712 frame_context.block_size_max,
706713 );
707 if (written_size > block_size_maximum) return error.BlockSizeOverMaximum;
708714 const written_slice = ring_buffer.sliceLast(written_size);
709715 try result.appendSlice(written_slice.first);
710716 try result.appendSlice(written_slice.second);
711 if (hasher_opt) |*hasher| {
717 if (frame_context.hasher_opt) |*hasher| {
712718 hasher.update(written_slice.first);
713719 hasher.update(written_slice.second);
714720 }
715721 if (block_header.last_block) break;
716722 }
717723
718 if (frame_header.descriptor.content_checksum_flag) {
724 if (frame_context.has_checksum) {
719725 const checksum = readIntSlice(u32, src[consumed_count .. consumed_count + 4]);
720726 consumed_count += 4;
721 if (hasher_opt) |*hasher| {
727 if (frame_context.hasher_opt) |*hasher| {
722728 if (checksum != computeChecksum(hasher)) return error.ChecksumFailure;
723729 }
724730 }
......@@ -741,9 +747,9 @@ pub fn decodeFrameBlocks(
741747 hash: ?*std.hash.XxHash64,
742748) DecodeBlockError!usize {
743749 // These tables take 7680 bytes
744 var literal_fse_data: [literal_table_size_max]Table.Fse = undefined;
745 var match_fse_data: [match_table_size_max]Table.Fse = undefined;
746 var offset_fse_data: [offset_table_size_max]Table.Fse = undefined;
750 var literal_fse_data: [types.compressed_block.table_size_max.literal]Table.Fse = undefined;
751 var match_fse_data: [types.compressed_block.table_size_max.match]Table.Fse = undefined;
752 var offset_fse_data: [types.compressed_block.table_size_max.offset]Table.Fse = undefined;
747753
748754 var block_header = decodeBlockHeader(src[0..3]);
749755 var bytes_read: usize = 3;
......@@ -766,6 +772,8 @@ pub fn decodeFrameBlocks(
766772 .fse_tables_undefined = true,
767773
768774 .literal_written_count = 0,
775 .literal_header = undefined,
776 .literal_streams = undefined,
769777 .literal_stream_reader = undefined,
770778 .literal_stream_index = undefined,
771779 .huffman_tree = null,
......@@ -867,7 +875,8 @@ pub fn decodeBlock(
867875 .compressed => {
868876 if (src.len < block_size) return error.MalformedBlockSize;
869877 var bytes_read: usize = 0;
870 const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock;
878 const literals = decodeLiteralsSection(src, &bytes_read) catch
879 return error.MalformedCompressedBlock;
871880 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
872881 return error.MalformedCompressedBlock;
873882
......@@ -889,7 +898,6 @@ pub fn decodeBlock(
889898 const decompressed_size = decode_state.decodeSequenceSlice(
890899 dest,
891900 write_pos,
892 literals,
893901 &bit_stream,
894902 sequence_size_limit,
895903 i == sequences_header.sequence_count - 1,
......@@ -903,12 +911,11 @@ pub fn decodeBlock(
903911
904912 if (decode_state.literal_written_count < literals.header.regenerated_size) {
905913 const len = literals.header.regenerated_size - decode_state.literal_written_count;
906 decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], literals, len) catch
914 decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], len) catch
907915 return error.MalformedCompressedBlock;
908916 bytes_written += len;
909917 }
910918
911 decode_state.literal_written_count = 0;
912919 assert(bytes_read == block_header.block_size);
913920 consumed_count.* += bytes_read;
914921 return bytes_written;
......@@ -936,7 +943,8 @@ pub fn decodeBlockRingBuffer(
936943 .compressed => {
937944 if (src.len < block_size) return error.MalformedBlockSize;
938945 var bytes_read: usize = 0;
939 const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock;
946 const literals = decodeLiteralsSection(src, &bytes_read) catch
947 return error.MalformedCompressedBlock;
940948 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
941949 return error.MalformedCompressedBlock;
942950
......@@ -956,7 +964,6 @@ pub fn decodeBlockRingBuffer(
956964 while (i < sequences_header.sequence_count) : (i += 1) {
957965 const decompressed_size = decode_state.decodeSequenceRingBuffer(
958966 dest,
959 literals,
960967 &bit_stream,
961968 sequence_size_limit,
962969 i == sequences_header.sequence_count - 1,
......@@ -970,14 +977,14 @@ pub fn decodeBlockRingBuffer(
970977
971978 if (decode_state.literal_written_count < literals.header.regenerated_size) {
972979 const len = literals.header.regenerated_size - decode_state.literal_written_count;
973 decode_state.decodeLiteralsRingBuffer(dest, literals, len) catch
980 decode_state.decodeLiteralsRingBuffer(dest, len) catch
974981 return error.MalformedCompressedBlock;
975982 bytes_written += len;
976983 }
977984
978 decode_state.literal_written_count = 0;
979985 assert(bytes_read == block_header.block_size);
980986 consumed_count.* += bytes_read;
987 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
981988 return bytes_written;
982989 },
983990 .reserved => return error.ReservedBlock,
lib/std/compress/zstandard/types.zig+5
......@@ -386,6 +386,11 @@ pub const compressed_block = struct {
386386 pub const match = 6;
387387 pub const offset = 5;
388388 };
389 pub const table_size_max = struct {
390 pub const literal = 1 << table_accuracy_log_max.literal;
391 pub const match = 1 << table_accuracy_log_max.match;
392 pub const offset = 1 << table_accuracy_log_max.match;
393 };
389394};
390395
391396test {