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 {...@@ -81,6 +81,8 @@ pub const DecodeState = struct {
8181
82 literal_stream_reader: ReverseBitReader,82 literal_stream_reader: ReverseBitReader,
83 literal_stream_index: usize,83 literal_stream_index: usize,
84 literal_streams: LiteralsSection.Streams,
85 literal_header: LiteralsSection.Header,
84 huffman_tree: ?LiteralsSection.HuffmanTree,86 huffman_tree: ?LiteralsSection.HuffmanTree,
8587
86 literal_written_count: usize,88 literal_written_count: usize,
...@@ -105,6 +107,10 @@ pub const DecodeState = struct {...@@ -105,6 +107,10 @@ pub const DecodeState = struct {
105 literals: LiteralsSection,107 literals: LiteralsSection,
106 sequences_header: SequencesSection.Header,108 sequences_header: SequencesSection.Header,
107 ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize {109 ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize {
110 self.literal_written_count = 0;
111 self.literal_header = literals.header;
112 self.literal_streams = literals.streams;
113
108 if (literals.huffman_tree) |tree| {114 if (literals.huffman_tree) |tree| {
109 self.huffman_tree = tree;115 self.huffman_tree = tree;
110 } else if (literals.header.block_type == .treeless and self.huffman_tree == null) {116 } else if (literals.header.block_type == .treeless and self.huffman_tree == null) {
...@@ -293,12 +299,11 @@ pub const DecodeState = struct {...@@ -293,12 +299,11 @@ pub const DecodeState = struct {
293 self: *DecodeState,299 self: *DecodeState,
294 dest: []u8,300 dest: []u8,
295 write_pos: usize,301 write_pos: usize,
296 literals: LiteralsSection,
297 sequence: Sequence,302 sequence: Sequence,
298 ) (error{MalformedSequence} || DecodeLiteralsError)!void {303 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
299 if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence;304 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);
302 const copy_start = write_pos + sequence.literal_length - sequence.offset;307 const copy_start = write_pos + sequence.literal_length - sequence.offset;
303 const copy_end = copy_start + sequence.match_length;308 const copy_end = copy_start + sequence.match_length;
304 // NOTE: we ignore the usage message for std.mem.copy and copy with dest.ptr >= src.ptr309 // 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 {...@@ -309,12 +314,11 @@ pub const DecodeState = struct {
309 fn executeSequenceRingBuffer(314 fn executeSequenceRingBuffer(
310 self: *DecodeState,315 self: *DecodeState,
311 dest: *RingBuffer,316 dest: *RingBuffer,
312 literals: LiteralsSection,
313 sequence: Sequence,317 sequence: Sequence,
314 ) (error{MalformedSequence} || DecodeLiteralsError)!void {318 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
315 if (sequence.offset > dest.data.len) return error.MalformedSequence;319 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);
318 const copy_start = dest.write_index + dest.data.len - sequence.offset;322 const copy_start = dest.write_index + dest.data.len - sequence.offset;
319 const copy_slice = dest.sliceAt(copy_start, sequence.match_length);323 const copy_slice = dest.sliceAt(copy_start, sequence.match_length);
320 // TODO: would std.mem.copy and figuring out dest slice be better/faster?324 // TODO: would std.mem.copy and figuring out dest slice be better/faster?
...@@ -328,6 +332,7 @@ pub const DecodeState = struct {...@@ -328,6 +332,7 @@ pub const DecodeState = struct {
328 MalformedSequence,332 MalformedSequence,
329 MalformedFseBits,333 MalformedFseBits,
330 } || DecodeLiteralsError;334 } || DecodeLiteralsError;
335
331 /// Decode one sequence from `bit_reader` into `dest`, written starting at336 /// Decode one sequence from `bit_reader` into `dest`, written starting at
332 /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns337 /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns
333 /// `error.MalformedSequence` error if the decompressed sequence would be longer338 /// `error.MalformedSequence` error if the decompressed sequence would be longer
...@@ -340,7 +345,6 @@ pub const DecodeState = struct {...@@ -340,7 +345,6 @@ pub const DecodeState = struct {
340 self: *DecodeState,345 self: *DecodeState,
341 dest: []u8,346 dest: []u8,
342 write_pos: usize,347 write_pos: usize,
343 literals: LiteralsSection,
344 bit_reader: *ReverseBitReader,348 bit_reader: *ReverseBitReader,
345 sequence_size_limit: usize,349 sequence_size_limit: usize,
346 last_sequence: bool,350 last_sequence: bool,
...@@ -349,7 +353,7 @@ pub const DecodeState = struct {...@@ -349,7 +353,7 @@ pub const DecodeState = struct {
349 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;353 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
350 if (sequence_length > sequence_size_limit) return error.MalformedSequence;354 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);
353 if (!last_sequence) {357 if (!last_sequence) {
354 try self.updateState(.literal, bit_reader);358 try self.updateState(.literal, bit_reader);
355 try self.updateState(.match, bit_reader);359 try self.updateState(.match, bit_reader);
...@@ -362,7 +366,6 @@ pub const DecodeState = struct {...@@ -362,7 +366,6 @@ pub const DecodeState = struct {
362 pub fn decodeSequenceRingBuffer(366 pub fn decodeSequenceRingBuffer(
363 self: *DecodeState,367 self: *DecodeState,
364 dest: *RingBuffer,368 dest: *RingBuffer,
365 literals: LiteralsSection,
366 bit_reader: anytype,369 bit_reader: anytype,
367 sequence_size_limit: usize,370 sequence_size_limit: usize,
368 last_sequence: bool,371 last_sequence: bool,
...@@ -371,7 +374,7 @@ pub const DecodeState = struct {...@@ -371,7 +374,7 @@ pub const DecodeState = struct {
371 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;374 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
372 if (sequence_length > sequence_size_limit) return error.MalformedSequence;375 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
373376
374 try self.executeSequenceRingBuffer(dest, literals, sequence);377 try self.executeSequenceRingBuffer(dest, sequence);
375 if (!last_sequence) {378 if (!last_sequence) {
376 try self.updateState(.literal, bit_reader);379 try self.updateState(.literal, bit_reader);
377 try self.updateState(.match, bit_reader);380 try self.updateState(.match, bit_reader);
...@@ -382,13 +385,12 @@ pub const DecodeState = struct {...@@ -382,13 +385,12 @@ pub const DecodeState = struct {
382385
383 fn nextLiteralMultiStream(386 fn nextLiteralMultiStream(
384 self: *DecodeState,387 self: *DecodeState,
385 literals: LiteralsSection,
386 ) error{BitStreamHasNoStartBit}!void {388 ) error{BitStreamHasNoStartBit}!void {
387 self.literal_stream_index += 1;389 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]);
389 }391 }
390392
391 fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void {393 pub fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
392 try self.literal_stream_reader.init(bytes);394 try self.literal_stream_reader.init(bytes);
393 }395 }
394396
...@@ -400,11 +402,10 @@ pub const DecodeState = struct {...@@ -400,11 +402,10 @@ pub const DecodeState = struct {
400 self: *DecodeState,402 self: *DecodeState,
401 comptime T: type,403 comptime T: type,
402 bit_count_to_read: usize,404 bit_count_to_read: usize,
403 literals: LiteralsSection,
404 ) LiteralBitsError!T {405 ) LiteralBitsError!T {
405 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {406 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 if (self.literal_streams == .four and self.literal_stream_index < 3) {
407 try self.nextLiteralMultiStream(literals);408 try self.nextLiteralMultiStream();
408 break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch409 break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch
409 return error.UnexpectedEndOfLiteralStream;410 return error.UnexpectedEndOfLiteralStream;
410 } else {411 } else {
...@@ -427,23 +428,22 @@ pub const DecodeState = struct {...@@ -427,23 +428,22 @@ pub const DecodeState = struct {
427 pub fn decodeLiteralsSlice(428 pub fn decodeLiteralsSlice(
428 self: *DecodeState,429 self: *DecodeState,
429 dest: []u8,430 dest: []u8,
430 literals: LiteralsSection,
431 len: usize,431 len: usize,
432 ) DecodeLiteralsError!void {432 ) 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)
434 return error.MalformedLiteralsLength;434 return error.MalformedLiteralsLength;
435435
436 switch (literals.header.block_type) {436 switch (self.literal_header.block_type) {
437 .raw => {437 .raw => {
438 const literals_end = self.literal_written_count + len;438 const literals_end = self.literal_written_count + len;
439 const literal_data = literals.streams.one[self.literal_written_count..literals_end];439 const literal_data = self.literal_streams.one[self.literal_written_count..literals_end];
440 std.mem.copy(u8, dest, literal_data);440 std.mem.copy(u8, dest, literal_data);
441 self.literal_written_count += len;441 self.literal_written_count += len;
442 },442 },
443 .rle => {443 .rle => {
444 var i: usize = 0;444 var i: usize = 0;
445 while (i < len) : (i += 1) {445 while (i < len) : (i += 1) {
446 dest[i] = literals.streams.one[0];446 dest[i] = self.literal_streams.one[0];
447 }447 }
448 self.literal_written_count += len;448 self.literal_written_count += len;
449 },449 },
...@@ -462,7 +462,7 @@ pub const DecodeState = struct {...@@ -462,7 +462,7 @@ pub const DecodeState = struct {
462 while (i < len) : (i += 1) {462 while (i < len) : (i += 1) {
463 var prefix: u16 = 0;463 var prefix: u16 = 0;
464 while (true) {464 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);
466 prefix <<= bit_count_to_read;466 prefix <<= bit_count_to_read;
467 prefix |= new_bits;467 prefix |= new_bits;
468 bits_read += bit_count_to_read;468 bits_read += bit_count_to_read;
...@@ -496,23 +496,22 @@ pub const DecodeState = struct {...@@ -496,23 +496,22 @@ pub const DecodeState = struct {
496 pub fn decodeLiteralsRingBuffer(496 pub fn decodeLiteralsRingBuffer(
497 self: *DecodeState,497 self: *DecodeState,
498 dest: *RingBuffer,498 dest: *RingBuffer,
499 literals: LiteralsSection,
500 len: usize,499 len: usize,
501 ) DecodeLiteralsError!void {500 ) 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)
503 return error.MalformedLiteralsLength;502 return error.MalformedLiteralsLength;
504503
505 switch (literals.header.block_type) {504 switch (self.literal_header.block_type) {
506 .raw => {505 .raw => {
507 const literals_end = self.literal_written_count + len;506 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];
509 dest.writeSliceAssumeCapacity(literal_data);508 dest.writeSliceAssumeCapacity(literal_data);
510 self.literal_written_count += len;509 self.literal_written_count += len;
511 },510 },
512 .rle => {511 .rle => {
513 var i: usize = 0;512 var i: usize = 0;
514 while (i < len) : (i += 1) {513 while (i < len) : (i += 1) {
515 dest.writeAssumeCapacity(literals.streams.one[0]);514 dest.writeAssumeCapacity(self.literal_streams.one[0]);
516 }515 }
517 self.literal_written_count += len;516 self.literal_written_count += len;
518 },517 },
...@@ -531,7 +530,7 @@ pub const DecodeState = struct {...@@ -531,7 +530,7 @@ pub const DecodeState = struct {
531 while (i < len) : (i += 1) {530 while (i < len) : (i += 1) {
532 var prefix: u16 = 0;531 var prefix: u16 = 0;
533 while (true) {532 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);
535 prefix <<= bit_count_to_read;534 prefix <<= bit_count_to_read;
536 prefix |= new_bits;535 prefix |= new_bits;
537 bits_read += bit_count_to_read;536 bits_read += bit_count_to_read;
...@@ -569,10 +568,6 @@ pub const DecodeState = struct {...@@ -569,10 +568,6 @@ pub const DecodeState = struct {
569 }568 }
570};569};
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
576pub fn computeChecksum(hasher: *std.hash.XxHash64) u32 {571pub fn computeChecksum(hasher: *std.hash.XxHash64) u32 {
577 const hash = hasher.final();572 const hash = hasher.final();
578 return @intCast(u32, hash & 0xFFFFFFFF);573 return @intCast(u32, hash & 0xFFFFFFFF);
...@@ -625,6 +620,31 @@ pub fn decodeZStandardFrame(...@@ -625,6 +620,31 @@ pub fn decodeZStandardFrame(
625 return ReadWriteCount{ .read_count = consumed_count, .write_count = written_count };620 return ReadWriteCount{ .read_count = consumed_count, .write_count = written_count };
626}621}
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
628/// Decode a Zstandard from from `src` and return the decompressed bytes; see648/// Decode a Zstandard from from `src` and return the decompressed bytes; see
629/// `decodeZStandardFrame()`. Returns `error.WindowSizeUnknown` if the frame649/// `decodeZStandardFrame()`. Returns `error.WindowSizeUnknown` if the frame
630/// does not declare its content size or a window descriptor (this indicates a650/// does not declare its content size or a window descriptor (this indicates a
...@@ -639,33 +659,18 @@ pub fn decodeZStandardFrameAlloc(...@@ -639,33 +659,18 @@ pub fn decodeZStandardFrameAlloc(
639 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);659 assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number);
640 var consumed_count: usize = 4;660 var consumed_count: usize = 4;
641661
642 const frame_header = try decodeZStandardHeader(src[consumed_count..], &consumed_count);662 var frame_context = context: {
643663 const frame_header = try decodeZStandardHeader(src[consumed_count..], &consumed_count);
644 if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported;664 break :context try FrameContext.init(frame_header, window_size_max, verify_checksum);
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,
663 };665 };
664666
667 var ring_buffer = try RingBuffer.init(allocator, frame_context.window_size);
668 defer ring_buffer.deinit(allocator);
669
665 // These tables take 7680 bytes670 // These tables take 7680 bytes
666 var literal_fse_data: [literal_table_size_max]Table.Fse = undefined;671 var literal_fse_data: [types.compressed_block.table_size_max.literal]Table.Fse = undefined;
667 var match_fse_data: [match_table_size_max]Table.Fse = undefined;672 var match_fse_data: [types.compressed_block.table_size_max.match]Table.Fse = undefined;
668 var offset_fse_data: [offset_table_size_max]Table.Fse = undefined;673 var offset_fse_data: [types.compressed_block.table_size_max.offset]Table.Fse = undefined;
669674
670 var block_header = decodeBlockHeader(src[consumed_count..][0..3]);675 var block_header = decodeBlockHeader(src[consumed_count..][0..3]);
671 consumed_count += 3;676 consumed_count += 3;
...@@ -687,6 +692,8 @@ pub fn decodeZStandardFrameAlloc(...@@ -687,6 +692,8 @@ pub fn decodeZStandardFrameAlloc(
687 .fse_tables_undefined = true,692 .fse_tables_undefined = true,
688693
689 .literal_written_count = 0,694 .literal_written_count = 0,
695 .literal_header = undefined,
696 .literal_streams = undefined,
690 .literal_stream_reader = undefined,697 .literal_stream_reader = undefined,
691 .literal_stream_index = undefined,698 .literal_stream_index = undefined,
692 .huffman_tree = null,699 .huffman_tree = null,
...@@ -695,30 +702,29 @@ pub fn decodeZStandardFrameAlloc(...@@ -695,30 +702,29 @@ pub fn decodeZStandardFrameAlloc(
695 block_header = decodeBlockHeader(src[consumed_count..][0..3]);702 block_header = decodeBlockHeader(src[consumed_count..][0..3]);
696 consumed_count += 3;703 consumed_count += 3;
697 }) {704 }) {
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;
699 const written_size = try decodeBlockRingBuffer(706 const written_size = try decodeBlockRingBuffer(
700 &ring_buffer,707 &ring_buffer,
701 src[consumed_count..],708 src[consumed_count..],
702 block_header,709 block_header,
703 &decode_state,710 &decode_state,
704 &consumed_count,711 &consumed_count,
705 block_size_maximum,712 frame_context.block_size_max,
706 );713 );
707 if (written_size > block_size_maximum) return error.BlockSizeOverMaximum;
708 const written_slice = ring_buffer.sliceLast(written_size);714 const written_slice = ring_buffer.sliceLast(written_size);
709 try result.appendSlice(written_slice.first);715 try result.appendSlice(written_slice.first);
710 try result.appendSlice(written_slice.second);716 try result.appendSlice(written_slice.second);
711 if (hasher_opt) |*hasher| {717 if (frame_context.hasher_opt) |*hasher| {
712 hasher.update(written_slice.first);718 hasher.update(written_slice.first);
713 hasher.update(written_slice.second);719 hasher.update(written_slice.second);
714 }720 }
715 if (block_header.last_block) break;721 if (block_header.last_block) break;
716 }722 }
717723
718 if (frame_header.descriptor.content_checksum_flag) {724 if (frame_context.has_checksum) {
719 const checksum = readIntSlice(u32, src[consumed_count .. consumed_count + 4]);725 const checksum = readIntSlice(u32, src[consumed_count .. consumed_count + 4]);
720 consumed_count += 4;726 consumed_count += 4;
721 if (hasher_opt) |*hasher| {727 if (frame_context.hasher_opt) |*hasher| {
722 if (checksum != computeChecksum(hasher)) return error.ChecksumFailure;728 if (checksum != computeChecksum(hasher)) return error.ChecksumFailure;
723 }729 }
724 }730 }
...@@ -741,9 +747,9 @@ pub fn decodeFrameBlocks(...@@ -741,9 +747,9 @@ pub fn decodeFrameBlocks(
741 hash: ?*std.hash.XxHash64,747 hash: ?*std.hash.XxHash64,
742) DecodeBlockError!usize {748) DecodeBlockError!usize {
743 // These tables take 7680 bytes749 // These tables take 7680 bytes
744 var literal_fse_data: [literal_table_size_max]Table.Fse = undefined;750 var literal_fse_data: [types.compressed_block.table_size_max.literal]Table.Fse = undefined;
745 var match_fse_data: [match_table_size_max]Table.Fse = undefined;751 var match_fse_data: [types.compressed_block.table_size_max.match]Table.Fse = undefined;
746 var offset_fse_data: [offset_table_size_max]Table.Fse = undefined;752 var offset_fse_data: [types.compressed_block.table_size_max.offset]Table.Fse = undefined;
747753
748 var block_header = decodeBlockHeader(src[0..3]);754 var block_header = decodeBlockHeader(src[0..3]);
749 var bytes_read: usize = 3;755 var bytes_read: usize = 3;
...@@ -766,6 +772,8 @@ pub fn decodeFrameBlocks(...@@ -766,6 +772,8 @@ pub fn decodeFrameBlocks(
766 .fse_tables_undefined = true,772 .fse_tables_undefined = true,
767773
768 .literal_written_count = 0,774 .literal_written_count = 0,
775 .literal_header = undefined,
776 .literal_streams = undefined,
769 .literal_stream_reader = undefined,777 .literal_stream_reader = undefined,
770 .literal_stream_index = undefined,778 .literal_stream_index = undefined,
771 .huffman_tree = null,779 .huffman_tree = null,
...@@ -867,7 +875,8 @@ pub fn decodeBlock(...@@ -867,7 +875,8 @@ pub fn decodeBlock(
867 .compressed => {875 .compressed => {
868 if (src.len < block_size) return error.MalformedBlockSize;876 if (src.len < block_size) return error.MalformedBlockSize;
869 var bytes_read: usize = 0;877 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;
871 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch880 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
872 return error.MalformedCompressedBlock;881 return error.MalformedCompressedBlock;
873882
...@@ -889,7 +898,6 @@ pub fn decodeBlock(...@@ -889,7 +898,6 @@ pub fn decodeBlock(
889 const decompressed_size = decode_state.decodeSequenceSlice(898 const decompressed_size = decode_state.decodeSequenceSlice(
890 dest,899 dest,
891 write_pos,900 write_pos,
892 literals,
893 &bit_stream,901 &bit_stream,
894 sequence_size_limit,902 sequence_size_limit,
895 i == sequences_header.sequence_count - 1,903 i == sequences_header.sequence_count - 1,
...@@ -903,12 +911,11 @@ pub fn decodeBlock(...@@ -903,12 +911,11 @@ pub fn decodeBlock(
903911
904 if (decode_state.literal_written_count < literals.header.regenerated_size) {912 if (decode_state.literal_written_count < literals.header.regenerated_size) {
905 const len = literals.header.regenerated_size - decode_state.literal_written_count;913 const len = literals.header.regenerated_size - decode_state.literal_written_count;
906 decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], literals, len) catch914 decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], len) catch
907 return error.MalformedCompressedBlock;915 return error.MalformedCompressedBlock;
908 bytes_written += len;916 bytes_written += len;
909 }917 }
910918
911 decode_state.literal_written_count = 0;
912 assert(bytes_read == block_header.block_size);919 assert(bytes_read == block_header.block_size);
913 consumed_count.* += bytes_read;920 consumed_count.* += bytes_read;
914 return bytes_written;921 return bytes_written;
...@@ -936,7 +943,8 @@ pub fn decodeBlockRingBuffer(...@@ -936,7 +943,8 @@ pub fn decodeBlockRingBuffer(
936 .compressed => {943 .compressed => {
937 if (src.len < block_size) return error.MalformedBlockSize;944 if (src.len < block_size) return error.MalformedBlockSize;
938 var bytes_read: usize = 0;945 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;
940 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch948 const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch
941 return error.MalformedCompressedBlock;949 return error.MalformedCompressedBlock;
942950
...@@ -956,7 +964,6 @@ pub fn decodeBlockRingBuffer(...@@ -956,7 +964,6 @@ pub fn decodeBlockRingBuffer(
956 while (i < sequences_header.sequence_count) : (i += 1) {964 while (i < sequences_header.sequence_count) : (i += 1) {
957 const decompressed_size = decode_state.decodeSequenceRingBuffer(965 const decompressed_size = decode_state.decodeSequenceRingBuffer(
958 dest,966 dest,
959 literals,
960 &bit_stream,967 &bit_stream,
961 sequence_size_limit,968 sequence_size_limit,
962 i == sequences_header.sequence_count - 1,969 i == sequences_header.sequence_count - 1,
...@@ -970,14 +977,14 @@ pub fn decodeBlockRingBuffer(...@@ -970,14 +977,14 @@ pub fn decodeBlockRingBuffer(
970977
971 if (decode_state.literal_written_count < literals.header.regenerated_size) {978 if (decode_state.literal_written_count < literals.header.regenerated_size) {
972 const len = literals.header.regenerated_size - decode_state.literal_written_count;979 const len = literals.header.regenerated_size - decode_state.literal_written_count;
973 decode_state.decodeLiteralsRingBuffer(dest, literals, len) catch980 decode_state.decodeLiteralsRingBuffer(dest, len) catch
974 return error.MalformedCompressedBlock;981 return error.MalformedCompressedBlock;
975 bytes_written += len;982 bytes_written += len;
976 }983 }
977984
978 decode_state.literal_written_count = 0;
979 assert(bytes_read == block_header.block_size);985 assert(bytes_read == block_header.block_size);
980 consumed_count.* += bytes_read;986 consumed_count.* += bytes_read;
987 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
981 return bytes_written;988 return bytes_written;
982 },989 },
983 .reserved => return error.ReservedBlock,990 .reserved => return error.ReservedBlock,
lib/std/compress/zstandard/types.zig+5
...@@ -386,6 +386,11 @@ pub const compressed_block = struct {...@@ -386,6 +386,11 @@ pub const compressed_block = struct {
386 pub const match = 6;386 pub const match = 6;
387 pub const offset = 5;387 pub const offset = 5;
388 };388 };
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 };
389};394};
390395
391test {396test {