authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-01-22 16:11:47+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:05+11:00
log05e63f241edb2199e91ce29c488e104dfb826935
tree9784750af821ddcc80364e841b58dbaf26136d4d
parent18091723d5afa8001e0fd71274dc4b74d601d0e1

std.compress.zstandard: add functions decoding into ring buffer

This supports decoding frames that do not declare the content size or decoding in a streaming fashion.

2 files changed, 272 insertions(+), 3 deletions(-)

lib/std/compress/zstandard/RingBuffer.zig created+81
......@@ -0,0 +1,81 @@
1//! This ring buffer stores read and write indices while being able to utilise the full
2//! backing slice by incrementing the indices modulo twice the slice's length and reducing
3//! indices modulo the slice's length on slice access. This means that the bit of information
4//! distinguishing whether the buffer is full or empty in an implementation utilising
5//! and extra flag is stored in difference of the indices.
6
7const assert = @import("std").debug.assert;
8
9const RingBuffer = @This();
10
11data: []u8,
12read_index: usize,
13write_index: usize,
14
15pub fn mask(self: RingBuffer, index: usize) usize {
16 return index % self.data.len;
17}
18
19pub fn mask2(self: RingBuffer, index: usize) usize {
20 return index % (2 * self.data.len);
21}
22
23pub fn write(self: *RingBuffer, byte: u8) !void {
24 if (self.isFull()) return error.Full;
25 self.writeAssumeCapacity(byte);
26}
27
28pub fn writeAssumeCapacity(self: *RingBuffer, byte: u8) void {
29 self.data[self.mask(self.write_index)] = byte;
30 self.write_index = self.mask2(self.write_index + 1);
31}
32
33pub fn writeSlice(self: *RingBuffer, bytes: []const u8) !void {
34 if (self.len() + bytes.len > self.data.len) return error.Full;
35 self.writeSliceAssumeCapacity(bytes);
36}
37
38pub fn writeSliceAssumeCapacity(self: *RingBuffer, bytes: []const u8) void {
39 for (bytes) |b| self.writeAssumeCapacity(b);
40}
41
42pub fn read(self: *RingBuffer) ?u8 {
43 if (self.isEmpty()) return null;
44 const byte = self.data[self.mask(self.read_index)];
45 self.read_index = self.mask2(self.read_index + 1);
46 return byte;
47}
48
49pub fn isEmpty(self: RingBuffer) bool {
50 return self.write_index == self.read_index;
51}
52
53pub fn isFull(self: RingBuffer) bool {
54 return self.mask2(self.write_index + self.data.len) == self.read_index;
55}
56
57pub fn len(self: RingBuffer) usize {
58 const adjusted_write_index = self.write_index + @boolToInt(self.write_index < self.read_index) * 2 * self.data.len;
59 return adjusted_write_index - self.read_index;
60}
61
62const Slice = struct {
63 first: []u8,
64 second: []u8,
65};
66
67pub fn sliceAt(self: RingBuffer, start_unmasked: usize, length: usize) Slice {
68 assert(length <= self.data.len);
69 const slice1_start = self.mask(start_unmasked);
70 const slice1_end = @min(self.data.len, slice1_start + length);
71 const slice1 = self.data[slice1_start..slice1_end];
72 const slice2 = self.data[0 .. length - slice1.len];
73 return Slice{
74 .first = slice1,
75 .second = slice2,
76 };
77}
78
79pub fn sliceLast(self: RingBuffer, length: usize) Slice {
80 return self.sliceAt(self.write_index + self.data.len - length, length);
81}
lib/std/compress/zstandard/decompress.zig+191-3
......@@ -6,6 +6,7 @@ const frame = types.frame;
66const Literals = types.compressed_block.Literals;
77const Sequences = types.compressed_block.Sequences;
88const Table = types.compressed_block.Table;
9const RingBuffer = @import("RingBuffer.zig");
910
1011const readInt = std.mem.readIntLittle;
1112const readIntSlice = std.mem.readIntSliceLittle;
......@@ -214,7 +215,7 @@ const DecodeState = struct {
214215 }
215216
216217 fn executeSequenceSlice(self: *DecodeState, dest: []u8, write_pos: usize, literals: Literals, sequence: Sequence) !void {
217 try self.decodeLiteralsInto(dest[write_pos..], literals, sequence.literal_length);
218 try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length);
218219
219220 // TODO: should we validate offset against max_window_size?
220221 assert(sequence.offset <= write_pos + sequence.literal_length);
......@@ -225,6 +226,15 @@ const DecodeState = struct {
225226 std.mem.copy(u8, dest[write_pos + sequence.literal_length ..], dest[copy_start..copy_end]);
226227 }
227228
229 fn executeSequenceRingBuffer(self: *DecodeState, dest: *RingBuffer, literals: Literals, sequence: Sequence) !void {
230 try self.decodeLiteralsRingBuffer(dest, literals, sequence.literal_length);
231 // TODO: check that ring buffer window is full enough for match copies
232 const copy_slice = dest.sliceAt(dest.write_index + dest.data.len - sequence.offset, sequence.match_length);
233 // TODO: would std.mem.copy and figuring out dest slice be better/faster?
234 for (copy_slice.first) |b| dest.writeAssumeCapacity(b);
235 for (copy_slice.second) |b| dest.writeAssumeCapacity(b);
236 }
237
228238 fn decodeSequenceSlice(
229239 self: *DecodeState,
230240 dest: []u8,
......@@ -246,6 +256,31 @@ const DecodeState = struct {
246256 return sequence.match_length + sequence.literal_length;
247257 }
248258
259 fn decodeSequenceRingBuffer(
260 self: *DecodeState,
261 dest: *RingBuffer,
262 literals: Literals,
263 bit_reader: anytype,
264 last_sequence: bool,
265 ) !usize {
266 const sequence = try self.nextSequence(bit_reader);
267 try self.executeSequenceRingBuffer(dest, literals, sequence);
268 if (std.options.log_level == .debug) {
269 const sequence_length = sequence.literal_length + sequence.match_length;
270 const written_slice = dest.sliceLast(sequence_length);
271 log.debug("sequence decompressed into '{x}{x}'", .{
272 std.fmt.fmtSliceHexUpper(written_slice.first),
273 std.fmt.fmtSliceHexUpper(written_slice.second),
274 });
275 }
276 if (!last_sequence) {
277 try self.updateState(.literal, bit_reader);
278 try self.updateState(.match, bit_reader);
279 try self.updateState(.offset, bit_reader);
280 }
281 return sequence.match_length + sequence.literal_length;
282 }
283
249284 fn nextLiteralMultiStream(self: *DecodeState, literals: Literals) !void {
250285 self.literal_stream_index += 1;
251286 try self.initLiteralStream(literals.streams.four[self.literal_stream_index]);
......@@ -258,7 +293,7 @@ const DecodeState = struct {
258293 while (0 == try self.literal_stream_reader.readBitsNoEof(u1, 1)) {}
259294 }
260295
261 fn decodeLiteralsInto(self: *DecodeState, dest: []u8, literals: Literals, len: usize) !void {
296 fn decodeLiteralsSlice(self: *DecodeState, dest: []u8, literals: Literals, len: usize) !void {
262297 if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength;
263298 switch (literals.header.block_type) {
264299 .raw => {
......@@ -327,6 +362,74 @@ const DecodeState = struct {
327362 }
328363 }
329364
365 fn decodeLiteralsRingBuffer(self: *DecodeState, dest: *RingBuffer, literals: Literals, len: usize) !void {
366 if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength;
367 switch (literals.header.block_type) {
368 .raw => {
369 const literal_data = literals.streams.one[self.literal_written_count .. self.literal_written_count + len];
370 dest.writeSliceAssumeCapacity(literal_data);
371 self.literal_written_count += len;
372 },
373 .rle => {
374 var i: usize = 0;
375 while (i < len) : (i += 1) {
376 dest.writeAssumeCapacity(literals.streams.one[0]);
377 }
378 self.literal_written_count += len;
379 },
380 .compressed, .treeless => {
381 // const written_bytes_per_stream = (literals.header.regenerated_size + 3) / 4;
382 const huffman_tree = self.huffman_tree orelse unreachable;
383 const max_bit_count = huffman_tree.max_bit_count;
384 const starting_bit_count = Literals.HuffmanTree.weightToBitCount(
385 huffman_tree.nodes[huffman_tree.symbol_count_minus_one].weight,
386 max_bit_count,
387 );
388 var bits_read: u4 = 0;
389 var huffman_tree_index: usize = huffman_tree.symbol_count_minus_one;
390 var bit_count_to_read: u4 = starting_bit_count;
391 var i: usize = 0;
392 while (i < len) : (i += 1) {
393 var prefix: u16 = 0;
394 while (true) {
395 const new_bits = self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |err|
396 switch (err) {
397 error.EndOfStream => if (literals.streams == .four and self.literal_stream_index < 3) bits: {
398 try self.nextLiteralMultiStream(literals);
399 break :bits try self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read);
400 } else {
401 return error.UnexpectedEndOfLiteralStream;
402 },
403 };
404 prefix <<= bit_count_to_read;
405 prefix |= new_bits;
406 bits_read += bit_count_to_read;
407 const result = try huffman_tree.query(huffman_tree_index, prefix);
408
409 switch (result) {
410 .symbol => |sym| {
411 dest.writeAssumeCapacity(sym);
412 bit_count_to_read = starting_bit_count;
413 bits_read = 0;
414 huffman_tree_index = huffman_tree.symbol_count_minus_one;
415 break;
416 },
417 .index => |index| {
418 huffman_tree_index = index;
419 const bit_count = Literals.HuffmanTree.weightToBitCount(
420 huffman_tree.nodes[index].weight,
421 max_bit_count,
422 );
423 bit_count_to_read = bit_count - bits_read;
424 },
425 }
426 }
427 }
428 self.literal_written_count += len;
429 },
430 }
431 }
432
330433 fn getCode(self: *DecodeState, comptime choice: DataType) u32 {
331434 return switch (@field(self, @tagName(choice)).table) {
332435 .rle => |value| value,
......@@ -437,6 +540,14 @@ fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:
437540 return block_size;
438541}
439542
543fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) usize {
544 log.debug("writing raw block - size {d}", .{block_size});
545 const data = src[0..block_size];
546 dest.writeSliceAssumeCapacity(data);
547 consumed_count.* += block_size;
548 return block_size;
549}
550
440551fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) usize {
441552 log.debug("writing rle block - '{x}'x{d}", .{ src[0], block_size });
442553 var write_pos: usize = 0;
......@@ -447,6 +558,16 @@ fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count:
447558 return block_size;
448559}
449560
561fn decodeRleBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) usize {
562 log.debug("writing rle block - '{x}'x{d}", .{ src[0], block_size });
563 var write_pos: usize = 0;
564 while (write_pos < block_size) : (write_pos += 1) {
565 dest.writeAssumeCapacity(src[0]);
566 }
567 consumed_count.* += 1;
568 return block_size;
569}
570
450571fn prepareDecodeState(
451572 decode_state: *DecodeState,
452573 src: []const u8,
......@@ -545,7 +666,7 @@ pub fn decodeBlock(
545666 if (decode_state.literal_written_count < literals.header.regenerated_size) {
546667 log.debug("decoding remaining literals", .{});
547668 const len = literals.header.regenerated_size - decode_state.literal_written_count;
548 try decode_state.decodeLiteralsInto(dest[written_count + bytes_written ..], literals, len);
669 try decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], literals, len);
549670 log.debug("remaining decoded literals at {d}: {}", .{
550671 written_count,
551672 std.fmt.fmtSliceHexUpper(dest[written_count .. written_count + len]),
......@@ -562,6 +683,73 @@ pub fn decodeBlock(
562683 }
563684}
564685
686pub fn decodeBlockRingBuffer(
687 dest: *RingBuffer,
688 src: []const u8,
689 block_header: frame.ZStandard.Block.Header,
690 decode_state: *DecodeState,
691 consumed_count: *usize,
692 block_size_maximum: usize,
693) !usize {
694 const block_size = block_header.block_size;
695 if (block_size_maximum < block_size) return error.BlockSizeOverMaximum;
696 // TODO: we probably want to enable safety for release-fast and release-small (or insert custom checks)
697 switch (block_header.block_type) {
698 .raw => return decodeRawBlockRingBuffer(dest, src, block_size, consumed_count),
699 .rle => return decodeRleBlockRingBuffer(dest, src, block_size, consumed_count),
700 .compressed => {
701 var bytes_read: usize = 0;
702 const literals = try decodeLiteralsSection(src, &bytes_read);
703 const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read);
704
705 bytes_read += try prepareDecodeState(decode_state, src[bytes_read..], literals, sequences_header);
706
707 var bytes_written: usize = 0;
708 if (sequences_header.sequence_count > 0) {
709 const bit_stream_bytes = src[bytes_read..block_size];
710 var reverse_byte_reader = reversedByteReader(bit_stream_bytes);
711 var bit_stream = reverseBitReader(reverse_byte_reader.reader());
712
713 while (0 == try bit_stream.readBitsNoEof(u1, 1)) {}
714 try decode_state.readInitialState(&bit_stream);
715
716 var i: usize = 0;
717 while (i < sequences_header.sequence_count) : (i += 1) {
718 log.debug("decoding sequence {d}", .{i});
719 const decompressed_size = try decode_state.decodeSequenceRingBuffer(
720 dest,
721 literals,
722 &bit_stream,
723 i == sequences_header.sequence_count - 1,
724 );
725 bytes_written += decompressed_size;
726 }
727
728 bytes_read += bit_stream_bytes.len;
729 }
730
731 if (decode_state.literal_written_count < literals.header.regenerated_size) {
732 log.debug("decoding remaining literals", .{});
733 const len = literals.header.regenerated_size - decode_state.literal_written_count;
734 try decode_state.decodeLiteralsRingBuffer(dest, literals, len);
735 const written_slice = dest.sliceLast(len);
736 log.debug("remaining decoded literals at {d}: {}{}", .{
737 bytes_written,
738 std.fmt.fmtSliceHexUpper(written_slice.first),
739 std.fmt.fmtSliceHexUpper(written_slice.second),
740 });
741 bytes_written += len;
742 }
743
744 decode_state.literal_written_count = 0;
745 assert(bytes_read == block_header.block_size);
746 consumed_count.* += bytes_read;
747 return bytes_written;
748 },
749 .reserved => return error.FrameContainsReservedBlock,
750 }
751}
752
565753pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {
566754 const magic = readInt(u32, src[0..4]);
567755 assert(isSkippableMagic(magic));