authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-02 18:44:01+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log7e2755646f5c9cab9973708a79c8aaa369d148e7
treecbc0fcc22dddfde378f4b249da2e7b5457f578a9
parent6e3e72884bdc1a2f9f3ae372716b50803565e696

std.compress.zstandard: split decompressor into multiple files


6 files changed, 1538 insertions(+), 1466 deletions(-)

lib/std/compress/zstandard.zig+30-10
......@@ -13,7 +13,7 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
1313
1414 allocator: Allocator,
1515 in_reader: ReaderType,
16 decode_state: decompress.DecodeState,
16 decode_state: decompress.block.DecodeState,
1717 frame_context: decompress.FrameContext,
1818 buffer: RingBuffer,
1919 last_block: bool,
......@@ -24,7 +24,7 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
2424 sequence_buffer: []u8,
2525 checksum: if (verify_checksum) ?u32 else void,
2626
27 pub const Error = ReaderType.Error || error{ MalformedBlock, MalformedFrame, EndOfStream };
27 pub const Error = ReaderType.Error || error{ MalformedBlock, MalformedFrame };
2828
2929 pub const Reader = std.io.Reader(*Self, Error, read);
3030
......@@ -34,21 +34,41 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
3434 .zstandard => {
3535 const frame_context = context: {
3636 const frame_header = try decompress.decodeZStandardHeader(source);
37 break :context try decompress.FrameContext.init(frame_header, window_size_max, verify_checksum);
37 break :context try decompress.FrameContext.init(
38 frame_header,
39 window_size_max,
40 verify_checksum,
41 );
3842 };
3943
40 const literal_fse_buffer = try allocator.alloc(types.compressed_block.Table.Fse, types.compressed_block.table_size_max.literal);
44 const literal_fse_buffer = try allocator.alloc(
45 types.compressed_block.Table.Fse,
46 types.compressed_block.table_size_max.literal,
47 );
4148 errdefer allocator.free(literal_fse_buffer);
42 const match_fse_buffer = try allocator.alloc(types.compressed_block.Table.Fse, types.compressed_block.table_size_max.match);
49
50 const match_fse_buffer = try allocator.alloc(
51 types.compressed_block.Table.Fse,
52 types.compressed_block.table_size_max.match,
53 );
4354 errdefer allocator.free(match_fse_buffer);
44 const offset_fse_buffer = try allocator.alloc(types.compressed_block.Table.Fse, types.compressed_block.table_size_max.offset);
55
56 const offset_fse_buffer = try allocator.alloc(
57 types.compressed_block.Table.Fse,
58 types.compressed_block.table_size_max.offset,
59 );
4560 errdefer allocator.free(offset_fse_buffer);
4661
47 const decode_state = decompress.DecodeState.init(literal_fse_buffer, match_fse_buffer, offset_fse_buffer);
62 const decode_state = decompress.block.DecodeState.init(
63 literal_fse_buffer,
64 match_fse_buffer,
65 offset_fse_buffer,
66 );
4867 const buffer = try RingBuffer.init(allocator, frame_context.window_size);
4968
5069 const literals_data = try allocator.alloc(u8, window_size_max);
5170 errdefer allocator.free(literals_data);
71
5272 const sequence_data = try allocator.alloc(u8, window_size_max);
5373 errdefer allocator.free(sequence_data);
5474
......@@ -87,10 +107,10 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
87107 if (buffer.len == 0) return 0;
88108
89109 if (self.buffer.isEmpty() and !self.last_block) {
90 const header_bytes = try self.in_reader.readBytesNoEof(3);
91 const block_header = decompress.decodeBlockHeader(&header_bytes);
110 const header_bytes = self.in_reader.readBytesNoEof(3) catch return error.MalformedFrame;
111 const block_header = decompress.block.decodeBlockHeader(&header_bytes);
92112
93 decompress.decodeBlockReader(
113 decompress.block.decodeBlockReader(
94114 &self.buffer,
95115 self.in_reader,
96116 block_header,
lib/std/compress/zstandard/decode/block.zig created+1051
......@@ -0,0 +1,1051 @@
1const std = @import("std");
2const assert = std.debug.assert;
3
4const types = @import("../types.zig");
5const frame = types.frame;
6const Table = types.compressed_block.Table;
7const LiteralsSection = types.compressed_block.LiteralsSection;
8const SequencesSection = types.compressed_block.SequencesSection;
9
10const huffman = @import("huffman.zig");
11
12const RingBuffer = @import("../RingBuffer.zig");
13
14const readers = @import("../readers.zig");
15
16const decodeFseTable = @import("fse.zig").decodeFseTable;
17
18const readInt = std.mem.readIntLittle;
19
20pub const Error = error{
21 BlockSizeOverMaximum,
22 MalformedBlockSize,
23 ReservedBlock,
24 MalformedRleBlock,
25 MalformedCompressedBlock,
26 EndOfStream,
27};
28
29pub const DecodeState = struct {
30 repeat_offsets: [3]u32,
31
32 offset: StateData(8),
33 match: StateData(9),
34 literal: StateData(9),
35
36 offset_fse_buffer: []Table.Fse,
37 match_fse_buffer: []Table.Fse,
38 literal_fse_buffer: []Table.Fse,
39
40 fse_tables_undefined: bool,
41
42 literal_stream_reader: readers.ReverseBitReader,
43 literal_stream_index: usize,
44 literal_streams: LiteralsSection.Streams,
45 literal_header: LiteralsSection.Header,
46 huffman_tree: ?LiteralsSection.HuffmanTree,
47
48 literal_written_count: usize,
49
50 fn StateData(comptime max_accuracy_log: comptime_int) type {
51 return struct {
52 state: State,
53 table: Table,
54 accuracy_log: u8,
55
56 const State = std.meta.Int(.unsigned, max_accuracy_log);
57 };
58 }
59
60 pub fn init(
61 literal_fse_buffer: []Table.Fse,
62 match_fse_buffer: []Table.Fse,
63 offset_fse_buffer: []Table.Fse,
64 ) DecodeState {
65 return DecodeState{
66 .repeat_offsets = .{
67 types.compressed_block.start_repeated_offset_1,
68 types.compressed_block.start_repeated_offset_2,
69 types.compressed_block.start_repeated_offset_3,
70 },
71
72 .offset = undefined,
73 .match = undefined,
74 .literal = undefined,
75
76 .literal_fse_buffer = literal_fse_buffer,
77 .match_fse_buffer = match_fse_buffer,
78 .offset_fse_buffer = offset_fse_buffer,
79
80 .fse_tables_undefined = true,
81
82 .literal_written_count = 0,
83 .literal_header = undefined,
84 .literal_streams = undefined,
85 .literal_stream_reader = undefined,
86 .literal_stream_index = undefined,
87 .huffman_tree = null,
88 };
89 }
90
91 /// Prepare the decoder to decode a compressed block. Loads the literals
92 /// stream and Huffman tree from `literals` and reads the FSE tables from
93 /// `source`.
94 ///
95 /// Errors:
96 /// - returns `error.BitStreamHasNoStartBit` if the (reversed) literal bitstream's
97 /// first byte does not have any bits set.
98 /// - returns `error.TreelessLiteralsFirst` `literals` is a treeless literals section
99 /// and the decode state does not have a Huffman tree from a previous block.
100 pub fn prepare(
101 self: *DecodeState,
102 source: anytype,
103 literals: LiteralsSection,
104 sequences_header: SequencesSection.Header,
105 ) !void {
106 self.literal_written_count = 0;
107 self.literal_header = literals.header;
108 self.literal_streams = literals.streams;
109
110 if (literals.huffman_tree) |tree| {
111 self.huffman_tree = tree;
112 } else if (literals.header.block_type == .treeless and self.huffman_tree == null) {
113 return error.TreelessLiteralsFirst;
114 }
115
116 switch (literals.header.block_type) {
117 .raw, .rle => {},
118 .compressed, .treeless => {
119 self.literal_stream_index = 0;
120 switch (literals.streams) {
121 .one => |slice| try self.initLiteralStream(slice),
122 .four => |streams| try self.initLiteralStream(streams[0]),
123 }
124 },
125 }
126
127 if (sequences_header.sequence_count > 0) {
128 try self.updateFseTable(source, .literal, sequences_header.literal_lengths);
129 try self.updateFseTable(source, .offset, sequences_header.offsets);
130 try self.updateFseTable(source, .match, sequences_header.match_lengths);
131 self.fse_tables_undefined = false;
132 }
133 }
134
135 /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream`
136 /// if `bit_reader` does not contain enough bits.
137 pub fn readInitialFseState(self: *DecodeState, bit_reader: *readers.ReverseBitReader) error{EndOfStream}!void {
138 self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log);
139 self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log);
140 self.match.state = try bit_reader.readBitsNoEof(u9, self.match.accuracy_log);
141 }
142
143 fn updateRepeatOffset(self: *DecodeState, offset: u32) void {
144 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[1]);
145 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[2]);
146 self.repeat_offsets[0] = offset;
147 }
148
149 fn useRepeatOffset(self: *DecodeState, index: usize) u32 {
150 if (index == 1)
151 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[1])
152 else if (index == 2) {
153 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[2]);
154 std.mem.swap(u32, &self.repeat_offsets[1], &self.repeat_offsets[2]);
155 }
156 return self.repeat_offsets[0];
157 }
158
159 const DataType = enum { offset, match, literal };
160
161 fn updateState(
162 self: *DecodeState,
163 comptime choice: DataType,
164 bit_reader: *readers.ReverseBitReader,
165 ) error{ MalformedFseBits, EndOfStream }!void {
166 switch (@field(self, @tagName(choice)).table) {
167 .rle => {},
168 .fse => |table| {
169 const data = table[@field(self, @tagName(choice)).state];
170 const T = @TypeOf(@field(self, @tagName(choice))).State;
171 const bits_summand = try bit_reader.readBitsNoEof(T, data.bits);
172 const next_state = std.math.cast(
173 @TypeOf(@field(self, @tagName(choice))).State,
174 data.baseline + bits_summand,
175 ) orelse return error.MalformedFseBits;
176 @field(self, @tagName(choice)).state = next_state;
177 },
178 }
179 }
180
181 const FseTableError = error{
182 MalformedFseTable,
183 MalformedAccuracyLog,
184 RepeatModeFirst,
185 EndOfStream,
186 };
187
188 fn updateFseTable(
189 self: *DecodeState,
190 source: anytype,
191 comptime choice: DataType,
192 mode: SequencesSection.Header.Mode,
193 ) !void {
194 const field_name = @tagName(choice);
195 switch (mode) {
196 .predefined => {
197 @field(self, field_name).accuracy_log =
198 @field(types.compressed_block.default_accuracy_log, field_name);
199
200 @field(self, field_name).table =
201 @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table");
202 },
203 .rle => {
204 @field(self, field_name).accuracy_log = 0;
205 @field(self, field_name).table = .{ .rle = try source.readByte() };
206 },
207 .fse => {
208 var bit_reader = readers.bitReader(source);
209
210 const table_size = try decodeFseTable(
211 &bit_reader,
212 @field(types.compressed_block.table_symbol_count_max, field_name),
213 @field(types.compressed_block.table_accuracy_log_max, field_name),
214 @field(self, field_name ++ "_fse_buffer"),
215 );
216 @field(self, field_name).table = .{
217 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],
218 };
219 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);
220 },
221 .repeat => if (self.fse_tables_undefined) return error.RepeatModeFirst,
222 }
223 }
224
225 const Sequence = struct {
226 literal_length: u32,
227 match_length: u32,
228 offset: u32,
229 };
230
231 fn nextSequence(
232 self: *DecodeState,
233 bit_reader: *readers.ReverseBitReader,
234 ) error{ OffsetCodeTooLarge, EndOfStream }!Sequence {
235 const raw_code = self.getCode(.offset);
236 const offset_code = std.math.cast(u5, raw_code) orelse {
237 return error.OffsetCodeTooLarge;
238 };
239 const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code);
240
241 const match_code = self.getCode(.match);
242 const match = types.compressed_block.match_length_code_table[match_code];
243 const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]);
244
245 const literal_code = self.getCode(.literal);
246 const literal = types.compressed_block.literals_length_code_table[literal_code];
247 const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]);
248
249 const offset = if (offset_value > 3) offset: {
250 const offset = offset_value - 3;
251 self.updateRepeatOffset(offset);
252 break :offset offset;
253 } else offset: {
254 if (literal_length == 0) {
255 if (offset_value == 3) {
256 const offset = self.repeat_offsets[0] - 1;
257 self.updateRepeatOffset(offset);
258 break :offset offset;
259 }
260 break :offset self.useRepeatOffset(offset_value);
261 }
262 break :offset self.useRepeatOffset(offset_value - 1);
263 };
264
265 return .{
266 .literal_length = literal_length,
267 .match_length = match_length,
268 .offset = offset,
269 };
270 }
271
272 fn executeSequenceSlice(
273 self: *DecodeState,
274 dest: []u8,
275 write_pos: usize,
276 sequence: Sequence,
277 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
278 if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence;
279
280 try self.decodeLiteralsSlice(dest[write_pos..], sequence.literal_length);
281 const copy_start = write_pos + sequence.literal_length - sequence.offset;
282 const copy_end = copy_start + sequence.match_length;
283 // NOTE: we ignore the usage message for std.mem.copy and copy with dest.ptr >= src.ptr
284 // to allow repeats
285 std.mem.copy(u8, dest[write_pos + sequence.literal_length ..], dest[copy_start..copy_end]);
286 }
287
288 fn executeSequenceRingBuffer(
289 self: *DecodeState,
290 dest: *RingBuffer,
291 sequence: Sequence,
292 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
293 if (sequence.offset > dest.data.len) return error.MalformedSequence;
294
295 try self.decodeLiteralsRingBuffer(dest, sequence.literal_length);
296 const copy_start = dest.write_index + dest.data.len - sequence.offset;
297 const copy_slice = dest.sliceAt(copy_start, sequence.match_length);
298 // TODO: would std.mem.copy and figuring out dest slice be better/faster?
299 for (copy_slice.first) |b| dest.writeAssumeCapacity(b);
300 for (copy_slice.second) |b| dest.writeAssumeCapacity(b);
301 }
302
303 const DecodeSequenceError = error{
304 OffsetCodeTooLarge,
305 EndOfStream,
306 MalformedSequence,
307 MalformedFseBits,
308 } || DecodeLiteralsError;
309
310 /// Decode one sequence from `bit_reader` into `dest`, written starting at
311 /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns
312 /// `error.MalformedSequence` error if the decompressed sequence would be longer
313 /// than `sequence_size_limit` or the sequence's offset is too large; returns
314 /// `error.EndOfStream` if `bit_reader` does not contain enough bits; returns
315 /// `error.UnexpectedEndOfLiteralStream` if the decoder state's literal streams
316 /// do not contain enough literals for the sequence (this may mean the literal
317 /// stream or the sequence is malformed).
318 pub fn decodeSequenceSlice(
319 self: *DecodeState,
320 dest: []u8,
321 write_pos: usize,
322 bit_reader: *readers.ReverseBitReader,
323 sequence_size_limit: usize,
324 last_sequence: bool,
325 ) DecodeSequenceError!usize {
326 const sequence = try self.nextSequence(bit_reader);
327 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
328 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
329
330 try self.executeSequenceSlice(dest, write_pos, sequence);
331 if (!last_sequence) {
332 try self.updateState(.literal, bit_reader);
333 try self.updateState(.match, bit_reader);
334 try self.updateState(.offset, bit_reader);
335 }
336 return sequence_length;
337 }
338
339 /// Decode one sequence from `bit_reader` into `dest`; see `decodeSequenceSlice`.
340 pub fn decodeSequenceRingBuffer(
341 self: *DecodeState,
342 dest: *RingBuffer,
343 bit_reader: anytype,
344 sequence_size_limit: usize,
345 last_sequence: bool,
346 ) DecodeSequenceError!usize {
347 const sequence = try self.nextSequence(bit_reader);
348 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
349 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
350
351 try self.executeSequenceRingBuffer(dest, sequence);
352 if (!last_sequence) {
353 try self.updateState(.literal, bit_reader);
354 try self.updateState(.match, bit_reader);
355 try self.updateState(.offset, bit_reader);
356 }
357 return sequence_length;
358 }
359
360 fn nextLiteralMultiStream(
361 self: *DecodeState,
362 ) error{BitStreamHasNoStartBit}!void {
363 self.literal_stream_index += 1;
364 try self.initLiteralStream(self.literal_streams.four[self.literal_stream_index]);
365 }
366
367 pub fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
368 try self.literal_stream_reader.init(bytes);
369 }
370
371 const LiteralBitsError = error{
372 BitStreamHasNoStartBit,
373 UnexpectedEndOfLiteralStream,
374 };
375 fn readLiteralsBits(
376 self: *DecodeState,
377 comptime T: type,
378 bit_count_to_read: usize,
379 ) LiteralBitsError!T {
380 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {
381 if (self.literal_streams == .four and self.literal_stream_index < 3) {
382 try self.nextLiteralMultiStream();
383 break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch
384 return error.UnexpectedEndOfLiteralStream;
385 } else {
386 return error.UnexpectedEndOfLiteralStream;
387 }
388 };
389 }
390
391 const DecodeLiteralsError = error{
392 MalformedLiteralsLength,
393 PrefixNotFound,
394 } || LiteralBitsError;
395
396 /// Decode `len` bytes of literals into `dest`. `literals` should be the
397 /// `LiteralsSection` that was passed to `prepare()`. Returns
398 /// `error.MalformedLiteralsLength` if the number of literal bytes decoded by
399 /// `self` plus `len` is greater than the regenerated size of `literals`.
400 /// Returns `error.UnexpectedEndOfLiteralStream` and `error.PrefixNotFound` if
401 /// there are problems decoding Huffman compressed literals.
402 pub fn decodeLiteralsSlice(
403 self: *DecodeState,
404 dest: []u8,
405 len: usize,
406 ) DecodeLiteralsError!void {
407 if (self.literal_written_count + len > self.literal_header.regenerated_size)
408 return error.MalformedLiteralsLength;
409
410 switch (self.literal_header.block_type) {
411 .raw => {
412 const literals_end = self.literal_written_count + len;
413 const literal_data = self.literal_streams.one[self.literal_written_count..literals_end];
414 std.mem.copy(u8, dest, literal_data);
415 self.literal_written_count += len;
416 },
417 .rle => {
418 var i: usize = 0;
419 while (i < len) : (i += 1) {
420 dest[i] = self.literal_streams.one[0];
421 }
422 self.literal_written_count += len;
423 },
424 .compressed, .treeless => {
425 // const written_bytes_per_stream = (literals.header.regenerated_size + 3) / 4;
426 const huffman_tree = self.huffman_tree orelse unreachable;
427 const max_bit_count = huffman_tree.max_bit_count;
428 const starting_bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
429 huffman_tree.nodes[huffman_tree.symbol_count_minus_one].weight,
430 max_bit_count,
431 );
432 var bits_read: u4 = 0;
433 var huffman_tree_index: usize = huffman_tree.symbol_count_minus_one;
434 var bit_count_to_read: u4 = starting_bit_count;
435 var i: usize = 0;
436 while (i < len) : (i += 1) {
437 var prefix: u16 = 0;
438 while (true) {
439 const new_bits = self.readLiteralsBits(u16, bit_count_to_read) catch |err| {
440 return err;
441 };
442 prefix <<= bit_count_to_read;
443 prefix |= new_bits;
444 bits_read += bit_count_to_read;
445 const result = huffman_tree.query(huffman_tree_index, prefix) catch |err| {
446 return err;
447 };
448
449 switch (result) {
450 .symbol => |sym| {
451 dest[i] = sym;
452 bit_count_to_read = starting_bit_count;
453 bits_read = 0;
454 huffman_tree_index = huffman_tree.symbol_count_minus_one;
455 break;
456 },
457 .index => |index| {
458 huffman_tree_index = index;
459 const bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
460 huffman_tree.nodes[index].weight,
461 max_bit_count,
462 );
463 bit_count_to_read = bit_count - bits_read;
464 },
465 }
466 }
467 }
468 self.literal_written_count += len;
469 },
470 }
471 }
472
473 /// Decode literals into `dest`; see `decodeLiteralsSlice()`.
474 pub fn decodeLiteralsRingBuffer(
475 self: *DecodeState,
476 dest: *RingBuffer,
477 len: usize,
478 ) DecodeLiteralsError!void {
479 if (self.literal_written_count + len > self.literal_header.regenerated_size)
480 return error.MalformedLiteralsLength;
481
482 switch (self.literal_header.block_type) {
483 .raw => {
484 const literals_end = self.literal_written_count + len;
485 const literal_data = self.literal_streams.one[self.literal_written_count..literals_end];
486 dest.writeSliceAssumeCapacity(literal_data);
487 self.literal_written_count += len;
488 },
489 .rle => {
490 var i: usize = 0;
491 while (i < len) : (i += 1) {
492 dest.writeAssumeCapacity(self.literal_streams.one[0]);
493 }
494 self.literal_written_count += len;
495 },
496 .compressed, .treeless => {
497 // const written_bytes_per_stream = (literals.header.regenerated_size + 3) / 4;
498 const huffman_tree = self.huffman_tree orelse unreachable;
499 const max_bit_count = huffman_tree.max_bit_count;
500 const starting_bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
501 huffman_tree.nodes[huffman_tree.symbol_count_minus_one].weight,
502 max_bit_count,
503 );
504 var bits_read: u4 = 0;
505 var huffman_tree_index: usize = huffman_tree.symbol_count_minus_one;
506 var bit_count_to_read: u4 = starting_bit_count;
507 var i: usize = 0;
508 while (i < len) : (i += 1) {
509 var prefix: u16 = 0;
510 while (true) {
511 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read);
512 prefix <<= bit_count_to_read;
513 prefix |= new_bits;
514 bits_read += bit_count_to_read;
515 const result = try huffman_tree.query(huffman_tree_index, prefix);
516
517 switch (result) {
518 .symbol => |sym| {
519 dest.writeAssumeCapacity(sym);
520 bit_count_to_read = starting_bit_count;
521 bits_read = 0;
522 huffman_tree_index = huffman_tree.symbol_count_minus_one;
523 break;
524 },
525 .index => |index| {
526 huffman_tree_index = index;
527 const bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
528 huffman_tree.nodes[index].weight,
529 max_bit_count,
530 );
531 bit_count_to_read = bit_count - bits_read;
532 },
533 }
534 }
535 }
536 self.literal_written_count += len;
537 },
538 }
539 }
540
541 fn getCode(self: *DecodeState, comptime choice: DataType) u32 {
542 return switch (@field(self, @tagName(choice)).table) {
543 .rle => |value| value,
544 .fse => |table| table[@field(self, @tagName(choice)).state].symbol,
545 };
546 }
547};
548
549/// Decode a single block from `src` into `dest`. The beginning of `src` must be
550/// the start of the block content (i.e. directly after the block header).
551/// Increments `consumed_count` by the number of bytes read from `src` to decode
552/// the block and returns the decompressed size of the block.
553///
554/// Errors returned:
555///
556/// - `error.BlockSizeOverMaximum` if block's size is larger than 1 << 17 or
557/// `dest[written_count..].len`
558/// - `error.MalformedBlockSize` if `src.len` is smaller than the block size
559/// and the block is a raw or compressed block
560/// - `error.ReservedBlock` if the block is a reserved block
561/// - `error.MalformedRleBlock` if the block is an RLE block and `src.len < 1`
562/// - `error.MalformedCompressedBlock` if there are errors decoding a
563/// compressed block
564/// - `error.EndOfStream` if the sequence bit stream ends unexpectedly
565pub fn decodeBlock(
566 dest: []u8,
567 src: []const u8,
568 block_header: frame.ZStandard.Block.Header,
569 decode_state: *DecodeState,
570 consumed_count: *usize,
571 written_count: usize,
572) Error!usize {
573 const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB
574 const block_size = block_header.block_size;
575 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
576 switch (block_header.block_type) {
577 .raw => {
578 if (src.len < block_size) return error.MalformedBlockSize;
579 const data = src[0..block_size];
580 std.mem.copy(u8, dest[written_count..], data);
581 consumed_count.* += block_size;
582 return block_size;
583 },
584 .rle => {
585 if (src.len < 1) return error.MalformedRleBlock;
586 var write_pos: usize = written_count;
587 while (write_pos < block_size + written_count) : (write_pos += 1) {
588 dest[write_pos] = src[0];
589 }
590 consumed_count.* += 1;
591 return block_size;
592 },
593 .compressed => {
594 if (src.len < block_size) return error.MalformedBlockSize;
595 var bytes_read: usize = 0;
596 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
597 return error.MalformedCompressedBlock;
598 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
599 const fbs_reader = fbs.reader();
600 const sequences_header = decodeSequencesHeader(fbs_reader) catch
601 return error.MalformedCompressedBlock;
602
603 decode_state.prepare(fbs_reader, literals, sequences_header) catch
604 return error.MalformedCompressedBlock;
605
606 bytes_read += fbs.pos;
607
608 var bytes_written: usize = 0;
609 if (sequences_header.sequence_count > 0) {
610 const bit_stream_bytes = src[bytes_read..block_size];
611 var bit_stream: readers.ReverseBitReader = undefined;
612 bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock;
613
614 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
615
616 var sequence_size_limit = block_size_max;
617 var i: usize = 0;
618 while (i < sequences_header.sequence_count) : (i += 1) {
619 const write_pos = written_count + bytes_written;
620 const decompressed_size = decode_state.decodeSequenceSlice(
621 dest,
622 write_pos,
623 &bit_stream,
624 sequence_size_limit,
625 i == sequences_header.sequence_count - 1,
626 ) catch return error.MalformedCompressedBlock;
627 bytes_written += decompressed_size;
628 sequence_size_limit -= decompressed_size;
629 }
630
631 bytes_read += bit_stream_bytes.len;
632 }
633 if (bytes_read != block_size) return error.MalformedCompressedBlock;
634
635 if (decode_state.literal_written_count < literals.header.regenerated_size) {
636 const len = literals.header.regenerated_size - decode_state.literal_written_count;
637 decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], len) catch
638 return error.MalformedCompressedBlock;
639 bytes_written += len;
640 }
641
642 consumed_count.* += bytes_read;
643 return bytes_written;
644 },
645 .reserved => return error.ReservedBlock,
646 }
647}
648
649/// Decode a single block from `src` into `dest`; see `decodeBlock()`. Returns
650/// the size of the decompressed block, which can be used with `dest.sliceLast()`
651/// to get the decompressed bytes. `error.BlockSizeOverMaximum` is returned if
652/// the block's compressed or decompressed size is larger than `block_size_max`.
653pub fn decodeBlockRingBuffer(
654 dest: *RingBuffer,
655 src: []const u8,
656 block_header: frame.ZStandard.Block.Header,
657 decode_state: *DecodeState,
658 consumed_count: *usize,
659 block_size_max: usize,
660) Error!usize {
661 const block_size = block_header.block_size;
662 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
663 switch (block_header.block_type) {
664 .raw => {
665 if (src.len < block_size) return error.MalformedBlockSize;
666 const data = src[0..block_size];
667 dest.writeSliceAssumeCapacity(data);
668 consumed_count.* += block_size;
669 return block_size;
670 },
671 .rle => {
672 if (src.len < 1) return error.MalformedRleBlock;
673 var write_pos: usize = 0;
674 while (write_pos < block_size) : (write_pos += 1) {
675 dest.writeAssumeCapacity(src[0]);
676 }
677 consumed_count.* += 1;
678 return block_size;
679 },
680 .compressed => {
681 if (src.len < block_size) return error.MalformedBlockSize;
682 var bytes_read: usize = 0;
683 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
684 return error.MalformedCompressedBlock;
685 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
686 const fbs_reader = fbs.reader();
687 const sequences_header = decodeSequencesHeader(fbs_reader) catch
688 return error.MalformedCompressedBlock;
689
690 decode_state.prepare(fbs_reader, literals, sequences_header) catch
691 return error.MalformedCompressedBlock;
692
693 bytes_read += fbs.pos;
694
695 var bytes_written: usize = 0;
696 if (sequences_header.sequence_count > 0) {
697 const bit_stream_bytes = src[bytes_read..block_size];
698 var bit_stream: readers.ReverseBitReader = undefined;
699 bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock;
700
701 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
702
703 var sequence_size_limit = block_size_max;
704 var i: usize = 0;
705 while (i < sequences_header.sequence_count) : (i += 1) {
706 const decompressed_size = decode_state.decodeSequenceRingBuffer(
707 dest,
708 &bit_stream,
709 sequence_size_limit,
710 i == sequences_header.sequence_count - 1,
711 ) catch return error.MalformedCompressedBlock;
712 bytes_written += decompressed_size;
713 sequence_size_limit -= decompressed_size;
714 }
715
716 bytes_read += bit_stream_bytes.len;
717 }
718 if (bytes_read != block_size) return error.MalformedCompressedBlock;
719
720 if (decode_state.literal_written_count < literals.header.regenerated_size) {
721 const len = literals.header.regenerated_size - decode_state.literal_written_count;
722 decode_state.decodeLiteralsRingBuffer(dest, len) catch
723 return error.MalformedCompressedBlock;
724 bytes_written += len;
725 }
726
727 consumed_count.* += bytes_read;
728 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
729 return bytes_written;
730 },
731 .reserved => return error.ReservedBlock,
732 }
733}
734
735/// Decode a single block from `source` into `dest`. Literal and sequence data
736/// from the block is copied into `literals_buffer` and `sequence_buffer`, which
737/// must be large enough or `error.LiteralsBufferTooSmall` and
738/// `error.SequenceBufferTooSmall` are returned (the maximum block size is an
739/// upper bound for the size of both buffers). See `decodeBlock`
740/// and `decodeBlockRingBuffer` for function that can decode a block without
741/// these extra copies.
742pub fn decodeBlockReader(
743 dest: *RingBuffer,
744 source: anytype,
745 block_header: frame.ZStandard.Block.Header,
746 decode_state: *DecodeState,
747 block_size_max: usize,
748 literals_buffer: []u8,
749 sequence_buffer: []u8,
750) !void {
751 const block_size = block_header.block_size;
752 var block_reader_limited = std.io.limitedReader(source, block_size);
753 const block_reader = block_reader_limited.reader();
754 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
755 switch (block_header.block_type) {
756 .raw => {
757 const slice = dest.sliceAt(dest.write_index, block_size);
758 try source.readNoEof(slice.first);
759 try source.readNoEof(slice.second);
760 dest.write_index = dest.mask2(dest.write_index + block_size);
761 },
762 .rle => {
763 const byte = try source.readByte();
764 var i: usize = 0;
765 while (i < block_size) : (i += 1) {
766 dest.writeAssumeCapacity(byte);
767 }
768 },
769 .compressed => {
770 const literals = try decodeLiteralsSection(block_reader, literals_buffer);
771 const sequences_header = try decodeSequencesHeader(block_reader);
772
773 try decode_state.prepare(block_reader, literals, sequences_header);
774
775 if (sequences_header.sequence_count > 0) {
776 if (sequence_buffer.len < block_reader_limited.bytes_left)
777 return error.SequenceBufferTooSmall;
778
779 const size = try block_reader.readAll(sequence_buffer);
780 var bit_stream: readers.ReverseBitReader = undefined;
781 try bit_stream.init(sequence_buffer[0..size]);
782
783 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
784
785 var sequence_size_limit = block_size_max;
786 var i: usize = 0;
787 while (i < sequences_header.sequence_count) : (i += 1) {
788 const decompressed_size = decode_state.decodeSequenceRingBuffer(
789 dest,
790 &bit_stream,
791 sequence_size_limit,
792 i == sequences_header.sequence_count - 1,
793 ) catch return error.MalformedCompressedBlock;
794 sequence_size_limit -= decompressed_size;
795 }
796 }
797
798 if (decode_state.literal_written_count < literals.header.regenerated_size) {
799 const len = literals.header.regenerated_size - decode_state.literal_written_count;
800 decode_state.decodeLiteralsRingBuffer(dest, len) catch
801 return error.MalformedCompressedBlock;
802 }
803
804 decode_state.literal_written_count = 0;
805 assert(block_reader.readByte() == error.EndOfStream);
806 },
807 .reserved => return error.ReservedBlock,
808 }
809}
810
811/// Decode the header of a block.
812pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {
813 const last_block = src[0] & 1 == 1;
814 const block_type = @intToEnum(frame.ZStandard.Block.Type, (src[0] & 0b110) >> 1);
815 const block_size = ((src[0] & 0b11111000) >> 3) + (@as(u21, src[1]) << 5) + (@as(u21, src[2]) << 13);
816 return .{
817 .last_block = last_block,
818 .block_type = block_type,
819 .block_size = block_size,
820 };
821}
822
823pub fn decodeBlockHeaderSlice(src: []const u8) error{EndOfStream}!frame.ZStandard.Block.Header {
824 if (src.len < 3) return error.EndOfStream;
825 return decodeBlockHeader(src[0..3]);
826}
827
828/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
829/// number of bytes the section uses.
830///
831/// Errors:
832/// - returns `error.MalformedLiteralsHeader` if the header is invalid
833/// - returns `error.MalformedLiteralsSection` if there are errors decoding
834pub fn decodeLiteralsSectionSlice(
835 src: []const u8,
836 consumed_count: *usize,
837) (error{ MalformedLiteralsHeader, MalformedLiteralsSection, EndOfStream } || huffman.Error)!LiteralsSection {
838 var bytes_read: usize = 0;
839 const header = header: {
840 var fbs = std.io.fixedBufferStream(src);
841 defer bytes_read = fbs.pos;
842 break :header decodeLiteralsHeader(fbs.reader()) catch return error.MalformedLiteralsHeader;
843 };
844 switch (header.block_type) {
845 .raw => {
846 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;
847 const stream = src[bytes_read .. bytes_read + header.regenerated_size];
848 consumed_count.* += header.regenerated_size + bytes_read;
849 return LiteralsSection{
850 .header = header,
851 .huffman_tree = null,
852 .streams = .{ .one = stream },
853 };
854 },
855 .rle => {
856 if (src.len < bytes_read + 1) return error.MalformedLiteralsSection;
857 const stream = src[bytes_read .. bytes_read + 1];
858 consumed_count.* += 1 + bytes_read;
859 return LiteralsSection{
860 .header = header,
861 .huffman_tree = null,
862 .streams = .{ .one = stream },
863 };
864 },
865 .compressed, .treeless => {
866 const huffman_tree_start = bytes_read;
867 const huffman_tree = if (header.block_type == .compressed)
868 try huffman.decodeHuffmanTreeSlice(src[bytes_read..], &bytes_read)
869 else
870 null;
871 const huffman_tree_size = bytes_read - huffman_tree_start;
872 const total_streams_size = @as(usize, header.compressed_size.?) - huffman_tree_size;
873
874 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;
875 const stream_data = src[bytes_read .. bytes_read + total_streams_size];
876
877 const streams = try decodeStreams(header.size_format, stream_data);
878 consumed_count.* += bytes_read + total_streams_size;
879 return LiteralsSection{
880 .header = header,
881 .huffman_tree = huffman_tree,
882 .streams = streams,
883 };
884 },
885 }
886}
887
888/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
889/// number of bytes the section uses.
890///
891/// Errors:
892/// - returns `error.MalformedLiteralsHeader` if the header is invalid
893/// - returns `error.MalformedLiteralsSection` if there are errors decoding
894pub fn decodeLiteralsSection(
895 source: anytype,
896 buffer: []u8,
897) !LiteralsSection {
898 const header = try decodeLiteralsHeader(source);
899 switch (header.block_type) {
900 .raw => {
901 try source.readNoEof(buffer[0..header.regenerated_size]);
902 return LiteralsSection{
903 .header = header,
904 .huffman_tree = null,
905 .streams = .{ .one = buffer },
906 };
907 },
908 .rle => {
909 buffer[0] = try source.readByte();
910 return LiteralsSection{
911 .header = header,
912 .huffman_tree = null,
913 .streams = .{ .one = buffer[0..1] },
914 };
915 },
916 .compressed, .treeless => {
917 var counting_reader = std.io.countingReader(source);
918 const huffman_tree = if (header.block_type == .compressed)
919 try huffman.decodeHuffmanTree(counting_reader.reader(), buffer)
920 else
921 null;
922 const huffman_tree_size = counting_reader.bytes_read;
923 const total_streams_size = @as(usize, header.compressed_size.?) - @intCast(usize, huffman_tree_size);
924
925 if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall;
926 try source.readNoEof(buffer[0..total_streams_size]);
927 const stream_data = buffer[0..total_streams_size];
928
929 const streams = try decodeStreams(header.size_format, stream_data);
930 return LiteralsSection{
931 .header = header,
932 .huffman_tree = huffman_tree,
933 .streams = streams,
934 };
935 },
936 }
937}
938
939fn decodeStreams(size_format: u2, stream_data: []const u8) !LiteralsSection.Streams {
940 if (size_format == 0) {
941 return .{ .one = stream_data };
942 }
943
944 if (stream_data.len < 6) return error.MalformedLiteralsSection;
945
946 const stream_1_length = @as(usize, readInt(u16, stream_data[0..2]));
947 const stream_2_length = @as(usize, readInt(u16, stream_data[2..4]));
948 const stream_3_length = @as(usize, readInt(u16, stream_data[4..6]));
949
950 const stream_1_start = 6;
951 const stream_2_start = stream_1_start + stream_1_length;
952 const stream_3_start = stream_2_start + stream_2_length;
953 const stream_4_start = stream_3_start + stream_3_length;
954
955 return .{ .four = .{
956 stream_data[stream_1_start .. stream_1_start + stream_1_length],
957 stream_data[stream_2_start .. stream_2_start + stream_2_length],
958 stream_data[stream_3_start .. stream_3_start + stream_3_length],
959 stream_data[stream_4_start..],
960 } };
961}
962
963/// Decode a literals section header.
964pub fn decodeLiteralsHeader(source: anytype) !LiteralsSection.Header {
965 const byte0 = try source.readByte();
966 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);
967 const size_format = @intCast(u2, (byte0 & 0b1100) >> 2);
968 var regenerated_size: u20 = undefined;
969 var compressed_size: ?u18 = null;
970 switch (block_type) {
971 .raw, .rle => {
972 switch (size_format) {
973 0, 2 => {
974 regenerated_size = byte0 >> 3;
975 },
976 1 => regenerated_size = (byte0 >> 4) + (@as(u20, try source.readByte()) << 4),
977 3 => regenerated_size = (byte0 >> 4) +
978 (@as(u20, try source.readByte()) << 4) +
979 (@as(u20, try source.readByte()) << 12),
980 }
981 },
982 .compressed, .treeless => {
983 const byte1 = try source.readByte();
984 const byte2 = try source.readByte();
985 switch (size_format) {
986 0, 1 => {
987 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);
988 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);
989 },
990 2 => {
991 const byte3 = try source.readByte();
992 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);
993 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);
994 },
995 3 => {
996 const byte3 = try source.readByte();
997 const byte4 = try source.readByte();
998 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);
999 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);
1000 },
1001 }
1002 },
1003 }
1004 return LiteralsSection.Header{
1005 .block_type = block_type,
1006 .size_format = size_format,
1007 .regenerated_size = regenerated_size,
1008 .compressed_size = compressed_size,
1009 };
1010}
1011
1012/// Decode a sequences section header.
1013///
1014/// Errors:
1015/// - returns `error.ReservedBitSet` is the reserved bit is set
1016/// - returns `error.MalformedSequencesHeader` if the header is invalid
1017pub fn decodeSequencesHeader(
1018 source: anytype,
1019) !SequencesSection.Header {
1020 var sequence_count: u24 = undefined;
1021
1022 const byte0 = try source.readByte();
1023 if (byte0 == 0) {
1024 return SequencesSection.Header{
1025 .sequence_count = 0,
1026 .offsets = undefined,
1027 .match_lengths = undefined,
1028 .literal_lengths = undefined,
1029 };
1030 } else if (byte0 < 128) {
1031 sequence_count = byte0;
1032 } else if (byte0 < 255) {
1033 sequence_count = (@as(u24, (byte0 - 128)) << 8) + try source.readByte();
1034 } else {
1035 sequence_count = (try source.readByte()) + (@as(u24, try source.readByte()) << 8) + 0x7F00;
1036 }
1037
1038 const compression_modes = try source.readByte();
1039
1040 const matches_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00001100) >> 2);
1041 const offsets_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00110000) >> 4);
1042 const literal_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b11000000) >> 6);
1043 if (compression_modes & 0b11 != 0) return error.ReservedBitSet;
1044
1045 return SequencesSection.Header{
1046 .sequence_count = sequence_count,
1047 .offsets = offsets_mode,
1048 .match_lengths = matches_mode,
1049 .literal_lengths = literal_mode,
1050 };
1051}
lib/std/compress/zstandard/decode/fse.zig created+154
......@@ -0,0 +1,154 @@
1const std = @import("std");
2const assert = std.debug.assert;
3
4const types = @import("../types.zig");
5const Table = types.compressed_block.Table;
6
7pub fn decodeFseTable(
8 bit_reader: anytype,
9 expected_symbol_count: usize,
10 max_accuracy_log: u4,
11 entries: []Table.Fse,
12) !usize {
13 const accuracy_log_biased = try bit_reader.readBitsNoEof(u4, 4);
14 if (accuracy_log_biased > max_accuracy_log -| 5) return error.MalformedAccuracyLog;
15 const accuracy_log = accuracy_log_biased + 5;
16
17 var values: [256]u16 = undefined;
18 var value_count: usize = 0;
19
20 const total_probability = @as(u16, 1) << accuracy_log;
21 var accumulated_probability: u16 = 0;
22
23 while (accumulated_probability < total_probability) {
24 // WARNING: The RFC in poorly worded, and would suggest std.math.log2_int_ceil is correct here,
25 // but power of two (remaining probabilities + 1) need max bits set to 1 more.
26 const max_bits = std.math.log2_int(u16, total_probability - accumulated_probability + 1) + 1;
27 const small = try bit_reader.readBitsNoEof(u16, max_bits - 1);
28
29 const cutoff = (@as(u16, 1) << max_bits) - 1 - (total_probability - accumulated_probability + 1);
30
31 const value = if (small < cutoff)
32 small
33 else value: {
34 const value_read = small + (try bit_reader.readBitsNoEof(u16, 1) << (max_bits - 1));
35 break :value if (value_read < @as(u16, 1) << (max_bits - 1))
36 value_read
37 else
38 value_read - cutoff;
39 };
40
41 accumulated_probability += if (value != 0) value - 1 else 1;
42
43 values[value_count] = value;
44 value_count += 1;
45
46 if (value == 1) {
47 while (true) {
48 const repeat_flag = try bit_reader.readBitsNoEof(u2, 2);
49 var i: usize = 0;
50 while (i < repeat_flag) : (i += 1) {
51 values[value_count] = 1;
52 value_count += 1;
53 }
54 if (repeat_flag < 3) break;
55 }
56 }
57 }
58 bit_reader.alignToByte();
59
60 if (value_count < 2) return error.MalformedFseTable;
61 if (accumulated_probability != total_probability) return error.MalformedFseTable;
62 if (value_count > expected_symbol_count) return error.MalformedFseTable;
63
64 const table_size = total_probability;
65
66 try buildFseTable(values[0..value_count], entries[0..table_size]);
67 return table_size;
68}
69
70fn buildFseTable(values: []const u16, entries: []Table.Fse) !void {
71 const total_probability = @intCast(u16, entries.len);
72 const accuracy_log = std.math.log2_int(u16, total_probability);
73 assert(total_probability <= 1 << 9);
74
75 var less_than_one_count: usize = 0;
76 for (values) |value, i| {
77 if (value == 0) {
78 entries[entries.len - 1 - less_than_one_count] = Table.Fse{
79 .symbol = @intCast(u8, i),
80 .baseline = 0,
81 .bits = accuracy_log,
82 };
83 less_than_one_count += 1;
84 }
85 }
86
87 var position: usize = 0;
88 var temp_states: [1 << 9]u16 = undefined;
89 for (values) |value, symbol| {
90 if (value == 0 or value == 1) continue;
91 const probability = value - 1;
92
93 const state_share_dividend = std.math.ceilPowerOfTwo(u16, probability) catch
94 return error.MalformedFseTable;
95 const share_size = @divExact(total_probability, state_share_dividend);
96 const double_state_count = state_share_dividend - probability;
97 const single_state_count = probability - double_state_count;
98 const share_size_log = std.math.log2_int(u16, share_size);
99
100 var i: u16 = 0;
101 while (i < probability) : (i += 1) {
102 temp_states[i] = @intCast(u16, position);
103 position += (entries.len >> 1) + (entries.len >> 3) + 3;
104 position &= entries.len - 1;
105 while (position >= entries.len - less_than_one_count) {
106 position += (entries.len >> 1) + (entries.len >> 3) + 3;
107 position &= entries.len - 1;
108 }
109 }
110 std.sort.sort(u16, temp_states[0..probability], {}, std.sort.asc(u16));
111 i = 0;
112 while (i < probability) : (i += 1) {
113 entries[temp_states[i]] = if (i < double_state_count) Table.Fse{
114 .symbol = @intCast(u8, symbol),
115 .bits = share_size_log + 1,
116 .baseline = single_state_count * share_size + i * 2 * share_size,
117 } else Table.Fse{
118 .symbol = @intCast(u8, symbol),
119 .bits = share_size_log,
120 .baseline = (i - double_state_count) * share_size,
121 };
122 }
123 }
124}
125
126test buildFseTable {
127 const literals_length_default_values = [36]u16{
128 5, 4, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2,
129 3, 3, 3, 3, 3, 3, 3, 3, 3, 4, 3, 2, 2, 2, 2, 2,
130 0, 0, 0, 0,
131 };
132
133 const match_lengths_default_values = [53]u16{
134 2, 5, 4, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2,
135 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
136 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 0, 0,
137 0, 0, 0, 0, 0,
138 };
139
140 const offset_codes_default_values = [29]u16{
141 2, 2, 2, 2, 2, 2, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2,
142 2, 2, 2, 2, 2, 2, 2, 2, 0, 0, 0, 0, 0,
143 };
144
145 var entries: [64]Table.Fse = undefined;
146 try buildFseTable(&literals_length_default_values, &entries);
147 try std.testing.expectEqualSlices(Table.Fse, types.compressed_block.predefined_literal_fse_table.fse, &entries);
148
149 try buildFseTable(&match_lengths_default_values, &entries);
150 try std.testing.expectEqualSlices(Table.Fse, types.compressed_block.predefined_match_fse_table.fse, &entries);
151
152 try buildFseTable(&offset_codes_default_values, entries[0..32]);
153 try std.testing.expectEqualSlices(Table.Fse, types.compressed_block.predefined_offset_fse_table.fse, entries[0..32]);
154}
lib/std/compress/zstandard/decode/huffman.zig created+212
......@@ -0,0 +1,212 @@
1const std = @import("std");
2
3const types = @import("../types.zig");
4const LiteralsSection = types.compressed_block.LiteralsSection;
5const Table = types.compressed_block.Table;
6
7const readers = @import("../readers.zig");
8
9const decodeFseTable = @import("fse.zig").decodeFseTable;
10
11pub const Error = error{
12 MalformedHuffmanTree,
13 MalformedFseTable,
14 MalformedAccuracyLog,
15 EndOfStream,
16};
17
18fn decodeFseHuffmanTree(source: anytype, compressed_size: usize, buffer: []u8, weights: *[256]u4) !usize {
19 var stream = std.io.limitedReader(source, compressed_size);
20 var bit_reader = readers.bitReader(stream.reader());
21
22 var entries: [1 << 6]Table.Fse = undefined;
23 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
24 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
25 error.EndOfStream => return error.MalformedFseTable,
26 };
27 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
28
29 const amount = try stream.reader().readAll(buffer);
30 var huff_bits: readers.ReverseBitReader = undefined;
31 huff_bits.init(buffer[0..amount]) catch return error.MalformedHuffmanTree;
32
33 return assignWeights(&huff_bits, accuracy_log, &entries, weights);
34}
35
36fn decodeFseHuffmanTreeSlice(src: []const u8, compressed_size: usize, weights: *[256]u4) !usize {
37 if (src.len < compressed_size) return error.MalformedHuffmanTree;
38 var stream = std.io.fixedBufferStream(src[0..compressed_size]);
39 var counting_reader = std.io.countingReader(stream.reader());
40 var bit_reader = readers.bitReader(counting_reader.reader());
41
42 var entries: [1 << 6]Table.Fse = undefined;
43 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
44 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
45 error.EndOfStream => return error.MalformedFseTable,
46 };
47 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
48
49 const start_index = std.math.cast(usize, counting_reader.bytes_read) orelse return error.MalformedHuffmanTree;
50 var huff_data = src[start_index..compressed_size];
51 var huff_bits: readers.ReverseBitReader = undefined;
52 huff_bits.init(huff_data) catch return error.MalformedHuffmanTree;
53
54 return assignWeights(&huff_bits, accuracy_log, &entries, weights);
55}
56
57fn assignWeights(huff_bits: *readers.ReverseBitReader, accuracy_log: usize, entries: *[1 << 6]Table.Fse, weights: *[256]u4) !usize {
58 var i: usize = 0;
59 var even_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
60 var odd_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
61
62 while (i < 255) {
63 const even_data = entries[even_state];
64 var read_bits: usize = 0;
65 const even_bits = huff_bits.readBits(u32, even_data.bits, &read_bits) catch unreachable;
66 weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree;
67 i += 1;
68 if (read_bits < even_data.bits) {
69 weights[i] = std.math.cast(u4, entries[odd_state].symbol) orelse return error.MalformedHuffmanTree;
70 i += 1;
71 break;
72 }
73 even_state = even_data.baseline + even_bits;
74
75 read_bits = 0;
76 const odd_data = entries[odd_state];
77 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;
78 weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree;
79 i += 1;
80 if (read_bits < odd_data.bits) {
81 if (i == 256) return error.MalformedHuffmanTree;
82 weights[i] = std.math.cast(u4, entries[even_state].symbol) orelse return error.MalformedHuffmanTree;
83 i += 1;
84 break;
85 }
86 odd_state = odd_data.baseline + odd_bits;
87 } else return error.MalformedHuffmanTree;
88
89 return i + 1; // stream contains all but the last symbol
90}
91
92fn decodeDirectHuffmanTree(source: anytype, encoded_symbol_count: usize, weights: *[256]u4) !usize {
93 const weights_byte_count = (encoded_symbol_count + 1) / 2;
94 var i: usize = 0;
95 while (i < weights_byte_count) : (i += 1) {
96 const byte = try source.readByte();
97 weights[2 * i] = @intCast(u4, byte >> 4);
98 weights[2 * i + 1] = @intCast(u4, byte & 0xF);
99 }
100 return encoded_symbol_count + 1;
101}
102
103fn assignSymbols(weight_sorted_prefixed_symbols: []LiteralsSection.HuffmanTree.PrefixedSymbol, weights: [256]u4) usize {
104 for (weight_sorted_prefixed_symbols) |_, i| {
105 weight_sorted_prefixed_symbols[i] = .{
106 .symbol = @intCast(u8, i),
107 .weight = undefined,
108 .prefix = undefined,
109 };
110 }
111
112 std.sort.sort(
113 LiteralsSection.HuffmanTree.PrefixedSymbol,
114 weight_sorted_prefixed_symbols,
115 weights,
116 lessThanByWeight,
117 );
118
119 var prefix: u16 = 0;
120 var prefixed_symbol_count: usize = 0;
121 var sorted_index: usize = 0;
122 const symbol_count = weight_sorted_prefixed_symbols.len;
123 while (sorted_index < symbol_count) {
124 var symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
125 const weight = weights[symbol];
126 if (weight == 0) {
127 sorted_index += 1;
128 continue;
129 }
130
131 while (sorted_index < symbol_count) : ({
132 sorted_index += 1;
133 prefixed_symbol_count += 1;
134 prefix += 1;
135 }) {
136 symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
137 if (weights[symbol] != weight) {
138 prefix = ((prefix - 1) >> (weights[symbol] - weight)) + 1;
139 break;
140 }
141 weight_sorted_prefixed_symbols[prefixed_symbol_count].symbol = symbol;
142 weight_sorted_prefixed_symbols[prefixed_symbol_count].prefix = prefix;
143 weight_sorted_prefixed_symbols[prefixed_symbol_count].weight = weight;
144 }
145 }
146 return prefixed_symbol_count;
147}
148
149fn buildHuffmanTree(weights: *[256]u4, symbol_count: usize) LiteralsSection.HuffmanTree {
150 var weight_power_sum: u16 = 0;
151 for (weights[0 .. symbol_count - 1]) |value| {
152 if (value > 0) {
153 weight_power_sum += @as(u16, 1) << (value - 1);
154 }
155 }
156
157 // advance to next power of two (even if weight_power_sum is a power of 2)
158 const max_number_of_bits = std.math.log2_int(u16, weight_power_sum) + 1;
159 const next_power_of_two = @as(u16, 1) << max_number_of_bits;
160 weights[symbol_count - 1] = std.math.log2_int(u16, next_power_of_two - weight_power_sum) + 1;
161
162 var weight_sorted_prefixed_symbols: [256]LiteralsSection.HuffmanTree.PrefixedSymbol = undefined;
163 const prefixed_symbol_count = assignSymbols(weight_sorted_prefixed_symbols[0..symbol_count], weights.*);
164 const tree = LiteralsSection.HuffmanTree{
165 .max_bit_count = max_number_of_bits,
166 .symbol_count_minus_one = @intCast(u8, prefixed_symbol_count - 1),
167 .nodes = weight_sorted_prefixed_symbols,
168 };
169 return tree;
170}
171
172pub fn decodeHuffmanTree(source: anytype, buffer: []u8) !LiteralsSection.HuffmanTree {
173 const header = try source.readByte();
174 var weights: [256]u4 = undefined;
175 const symbol_count = if (header < 128)
176 // FSE compressed weights
177 try decodeFseHuffmanTree(source, header, buffer, &weights)
178 else
179 try decodeDirectHuffmanTree(source, header - 127, &weights);
180
181 return buildHuffmanTree(&weights, symbol_count);
182}
183
184pub fn decodeHuffmanTreeSlice(src: []const u8, consumed_count: *usize) Error!LiteralsSection.HuffmanTree {
185 if (src.len == 0) return error.MalformedHuffmanTree;
186 const header = src[0];
187 var bytes_read: usize = 1;
188 var weights: [256]u4 = undefined;
189 const symbol_count = if (header < 128) count: {
190 // FSE compressed weights
191 bytes_read += header;
192 break :count try decodeFseHuffmanTreeSlice(src[1..], header, &weights);
193 } else count: {
194 var fbs = std.io.fixedBufferStream(src[1..]);
195 defer bytes_read += fbs.pos;
196 break :count try decodeDirectHuffmanTree(fbs.reader(), header - 127, &weights);
197 };
198
199 consumed_count.* += bytes_read;
200 return buildHuffmanTree(&weights, symbol_count);
201}
202
203fn lessThanByWeight(
204 weights: [256]u4,
205 lhs: LiteralsSection.HuffmanTree.PrefixedSymbol,
206 rhs: LiteralsSection.HuffmanTree.PrefixedSymbol,
207) bool {
208 // NOTE: this function relies on the use of a stable sorting algorithm,
209 // otherwise a special case of if (weights[lhs] == weights[rhs]) return lhs < rhs;
210 // should be added
211 return weights[lhs.symbol] < weights[rhs.symbol];
212}
lib/std/compress/zstandard/decompress.zig+16-1456
......@@ -6,8 +6,13 @@ const frame = types.frame;
66const LiteralsSection = types.compressed_block.LiteralsSection;
77const SequencesSection = types.compressed_block.SequencesSection;
88const Table = types.compressed_block.Table;
9
10pub const block = @import("decode/block.zig");
11
912pub const RingBuffer = @import("RingBuffer.zig");
1013
14const readers = @import("readers.zig");
15
1116const readInt = std.mem.readIntLittle;
1217const readIntSlice = std.mem.readIntSliceLittle;
1318fn readVarInt(comptime T: type, bytes: []const u8) T {
......@@ -61,526 +66,6 @@ pub fn decodeFrame(
6166 };
6267}
6368
64pub const DecodeState = struct {
65 repeat_offsets: [3]u32,
66
67 offset: StateData(8),
68 match: StateData(9),
69 literal: StateData(9),
70
71 offset_fse_buffer: []Table.Fse,
72 match_fse_buffer: []Table.Fse,
73 literal_fse_buffer: []Table.Fse,
74
75 fse_tables_undefined: bool,
76
77 literal_stream_reader: ReverseBitReader,
78 literal_stream_index: usize,
79 literal_streams: LiteralsSection.Streams,
80 literal_header: LiteralsSection.Header,
81 huffman_tree: ?LiteralsSection.HuffmanTree,
82
83 literal_written_count: usize,
84
85 fn StateData(comptime max_accuracy_log: comptime_int) type {
86 return struct {
87 state: State,
88 table: Table,
89 accuracy_log: u8,
90
91 const State = std.meta.Int(.unsigned, max_accuracy_log);
92 };
93 }
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
126 /// Prepare the decoder to decode a compressed block. Loads the literals
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.
135 pub fn prepare(
136 self: *DecodeState,
137 source: anytype,
138 literals: LiteralsSection,
139 sequences_header: SequencesSection.Header,
140 ) !void {
141 self.literal_written_count = 0;
142 self.literal_header = literals.header;
143 self.literal_streams = literals.streams;
144
145 if (literals.huffman_tree) |tree| {
146 self.huffman_tree = tree;
147 } else if (literals.header.block_type == .treeless and self.huffman_tree == null) {
148 return error.TreelessLiteralsFirst;
149 }
150
151 switch (literals.header.block_type) {
152 .raw, .rle => {},
153 .compressed, .treeless => {
154 self.literal_stream_index = 0;
155 switch (literals.streams) {
156 .one => |slice| try self.initLiteralStream(slice),
157 .four => |streams| try self.initLiteralStream(streams[0]),
158 }
159 },
160 }
161
162 if (sequences_header.sequence_count > 0) {
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);
166 self.fse_tables_undefined = false;
167 }
168 }
169
170 /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream`
171 /// if `bit_reader` does not contain enough bits.
172 pub fn readInitialFseState(self: *DecodeState, bit_reader: *ReverseBitReader) error{EndOfStream}!void {
173 self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log);
174 self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log);
175 self.match.state = try bit_reader.readBitsNoEof(u9, self.match.accuracy_log);
176 }
177
178 fn updateRepeatOffset(self: *DecodeState, offset: u32) void {
179 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[1]);
180 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[2]);
181 self.repeat_offsets[0] = offset;
182 }
183
184 fn useRepeatOffset(self: *DecodeState, index: usize) u32 {
185 if (index == 1)
186 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[1])
187 else if (index == 2) {
188 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[2]);
189 std.mem.swap(u32, &self.repeat_offsets[1], &self.repeat_offsets[2]);
190 }
191 return self.repeat_offsets[0];
192 }
193
194 const DataType = enum { offset, match, literal };
195
196 fn updateState(
197 self: *DecodeState,
198 comptime choice: DataType,
199 bit_reader: *ReverseBitReader,
200 ) error{ MalformedFseBits, EndOfStream }!void {
201 switch (@field(self, @tagName(choice)).table) {
202 .rle => {},
203 .fse => |table| {
204 const data = table[@field(self, @tagName(choice)).state];
205 const T = @TypeOf(@field(self, @tagName(choice))).State;
206 const bits_summand = try bit_reader.readBitsNoEof(T, data.bits);
207 const next_state = std.math.cast(
208 @TypeOf(@field(self, @tagName(choice))).State,
209 data.baseline + bits_summand,
210 ) orelse return error.MalformedFseBits;
211 @field(self, @tagName(choice)).state = next_state;
212 },
213 }
214 }
215
216 const FseTableError = error{
217 MalformedFseTable,
218 MalformedAccuracyLog,
219 RepeatModeFirst,
220 EndOfStream,
221 };
222
223 fn updateFseTable(
224 self: *DecodeState,
225 source: anytype,
226 comptime choice: DataType,
227 mode: SequencesSection.Header.Mode,
228 ) !void {
229 const field_name = @tagName(choice);
230 switch (mode) {
231 .predefined => {
232 @field(self, field_name).accuracy_log =
233 @field(types.compressed_block.default_accuracy_log, field_name);
234
235 @field(self, field_name).table =
236 @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table");
237 },
238 .rle => {
239 @field(self, field_name).accuracy_log = 0;
240 @field(self, field_name).table = .{ .rle = try source.readByte() };
241 },
242 .fse => {
243 var bit_reader = bitReader(source);
244
245 const table_size = try decodeFseTable(
246 &bit_reader,
247 @field(types.compressed_block.table_symbol_count_max, field_name),
248 @field(types.compressed_block.table_accuracy_log_max, field_name),
249 @field(self, field_name ++ "_fse_buffer"),
250 );
251 @field(self, field_name).table = .{
252 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],
253 };
254 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);
255 },
256 .repeat => if (self.fse_tables_undefined) return error.RepeatModeFirst,
257 }
258 }
259
260 const Sequence = struct {
261 literal_length: u32,
262 match_length: u32,
263 offset: u32,
264 };
265
266 fn nextSequence(
267 self: *DecodeState,
268 bit_reader: *ReverseBitReader,
269 ) error{ OffsetCodeTooLarge, EndOfStream }!Sequence {
270 const raw_code = self.getCode(.offset);
271 const offset_code = std.math.cast(u5, raw_code) orelse {
272 return error.OffsetCodeTooLarge;
273 };
274 const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code);
275
276 const match_code = self.getCode(.match);
277 const match = types.compressed_block.match_length_code_table[match_code];
278 const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]);
279
280 const literal_code = self.getCode(.literal);
281 const literal = types.compressed_block.literals_length_code_table[literal_code];
282 const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]);
283
284 const offset = if (offset_value > 3) offset: {
285 const offset = offset_value - 3;
286 self.updateRepeatOffset(offset);
287 break :offset offset;
288 } else offset: {
289 if (literal_length == 0) {
290 if (offset_value == 3) {
291 const offset = self.repeat_offsets[0] - 1;
292 self.updateRepeatOffset(offset);
293 break :offset offset;
294 }
295 break :offset self.useRepeatOffset(offset_value);
296 }
297 break :offset self.useRepeatOffset(offset_value - 1);
298 };
299
300 return .{
301 .literal_length = literal_length,
302 .match_length = match_length,
303 .offset = offset,
304 };
305 }
306
307 fn executeSequenceSlice(
308 self: *DecodeState,
309 dest: []u8,
310 write_pos: usize,
311 sequence: Sequence,
312 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
313 if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence;
314
315 try self.decodeLiteralsSlice(dest[write_pos..], sequence.literal_length);
316 const copy_start = write_pos + sequence.literal_length - sequence.offset;
317 const copy_end = copy_start + sequence.match_length;
318 // NOTE: we ignore the usage message for std.mem.copy and copy with dest.ptr >= src.ptr
319 // to allow repeats
320 std.mem.copy(u8, dest[write_pos + sequence.literal_length ..], dest[copy_start..copy_end]);
321 }
322
323 fn executeSequenceRingBuffer(
324 self: *DecodeState,
325 dest: *RingBuffer,
326 sequence: Sequence,
327 ) (error{MalformedSequence} || DecodeLiteralsError)!void {
328 if (sequence.offset > dest.data.len) return error.MalformedSequence;
329
330 try self.decodeLiteralsRingBuffer(dest, sequence.literal_length);
331 const copy_start = dest.write_index + dest.data.len - sequence.offset;
332 const copy_slice = dest.sliceAt(copy_start, sequence.match_length);
333 // TODO: would std.mem.copy and figuring out dest slice be better/faster?
334 for (copy_slice.first) |b| dest.writeAssumeCapacity(b);
335 for (copy_slice.second) |b| dest.writeAssumeCapacity(b);
336 }
337
338 const DecodeSequenceError = error{
339 OffsetCodeTooLarge,
340 EndOfStream,
341 MalformedSequence,
342 MalformedFseBits,
343 } || DecodeLiteralsError;
344
345 /// Decode one sequence from `bit_reader` into `dest`, written starting at
346 /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns
347 /// `error.MalformedSequence` error if the decompressed sequence would be longer
348 /// than `sequence_size_limit` or the sequence's offset is too large; returns
349 /// `error.EndOfStream` if `bit_reader` does not contain enough bits; returns
350 /// `error.UnexpectedEndOfLiteralStream` if the decoder state's literal streams
351 /// do not contain enough literals for the sequence (this may mean the literal
352 /// stream or the sequence is malformed).
353 pub fn decodeSequenceSlice(
354 self: *DecodeState,
355 dest: []u8,
356 write_pos: usize,
357 bit_reader: *ReverseBitReader,
358 sequence_size_limit: usize,
359 last_sequence: bool,
360 ) DecodeSequenceError!usize {
361 const sequence = try self.nextSequence(bit_reader);
362 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
363 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
364
365 try self.executeSequenceSlice(dest, write_pos, sequence);
366 if (!last_sequence) {
367 try self.updateState(.literal, bit_reader);
368 try self.updateState(.match, bit_reader);
369 try self.updateState(.offset, bit_reader);
370 }
371 return sequence_length;
372 }
373
374 /// Decode one sequence from `bit_reader` into `dest`; see `decodeSequenceSlice`.
375 pub fn decodeSequenceRingBuffer(
376 self: *DecodeState,
377 dest: *RingBuffer,
378 bit_reader: anytype,
379 sequence_size_limit: usize,
380 last_sequence: bool,
381 ) DecodeSequenceError!usize {
382 const sequence = try self.nextSequence(bit_reader);
383 const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length;
384 if (sequence_length > sequence_size_limit) return error.MalformedSequence;
385
386 try self.executeSequenceRingBuffer(dest, sequence);
387 if (!last_sequence) {
388 try self.updateState(.literal, bit_reader);
389 try self.updateState(.match, bit_reader);
390 try self.updateState(.offset, bit_reader);
391 }
392 return sequence_length;
393 }
394
395 fn nextLiteralMultiStream(
396 self: *DecodeState,
397 ) error{BitStreamHasNoStartBit}!void {
398 self.literal_stream_index += 1;
399 try self.initLiteralStream(self.literal_streams.four[self.literal_stream_index]);
400 }
401
402 pub fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
403 try self.literal_stream_reader.init(bytes);
404 }
405
406 const LiteralBitsError = error{
407 BitStreamHasNoStartBit,
408 UnexpectedEndOfLiteralStream,
409 };
410 fn readLiteralsBits(
411 self: *DecodeState,
412 comptime T: type,
413 bit_count_to_read: usize,
414 ) LiteralBitsError!T {
415 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {
416 if (self.literal_streams == .four and self.literal_stream_index < 3) {
417 try self.nextLiteralMultiStream();
418 break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch
419 return error.UnexpectedEndOfLiteralStream;
420 } else {
421 return error.UnexpectedEndOfLiteralStream;
422 }
423 };
424 }
425
426 const DecodeLiteralsError = error{
427 MalformedLiteralsLength,
428 PrefixNotFound,
429 } || LiteralBitsError;
430
431 /// Decode `len` bytes of literals into `dest`. `literals` should be the
432 /// `LiteralsSection` that was passed to `prepare()`. Returns
433 /// `error.MalformedLiteralsLength` if the number of literal bytes decoded by
434 /// `self` plus `len` is greater than the regenerated size of `literals`.
435 /// Returns `error.UnexpectedEndOfLiteralStream` and `error.PrefixNotFound` if
436 /// there are problems decoding Huffman compressed literals.
437 pub fn decodeLiteralsSlice(
438 self: *DecodeState,
439 dest: []u8,
440 len: usize,
441 ) DecodeLiteralsError!void {
442 if (self.literal_written_count + len > self.literal_header.regenerated_size)
443 return error.MalformedLiteralsLength;
444
445 switch (self.literal_header.block_type) {
446 .raw => {
447 const literals_end = self.literal_written_count + len;
448 const literal_data = self.literal_streams.one[self.literal_written_count..literals_end];
449 std.mem.copy(u8, dest, literal_data);
450 self.literal_written_count += len;
451 },
452 .rle => {
453 var i: usize = 0;
454 while (i < len) : (i += 1) {
455 dest[i] = self.literal_streams.one[0];
456 }
457 self.literal_written_count += len;
458 },
459 .compressed, .treeless => {
460 // const written_bytes_per_stream = (literals.header.regenerated_size + 3) / 4;
461 const huffman_tree = self.huffman_tree orelse unreachable;
462 const max_bit_count = huffman_tree.max_bit_count;
463 const starting_bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
464 huffman_tree.nodes[huffman_tree.symbol_count_minus_one].weight,
465 max_bit_count,
466 );
467 var bits_read: u4 = 0;
468 var huffman_tree_index: usize = huffman_tree.symbol_count_minus_one;
469 var bit_count_to_read: u4 = starting_bit_count;
470 var i: usize = 0;
471 while (i < len) : (i += 1) {
472 var prefix: u16 = 0;
473 while (true) {
474 const new_bits = self.readLiteralsBits(u16, bit_count_to_read) catch |err| {
475 return err;
476 };
477 prefix <<= bit_count_to_read;
478 prefix |= new_bits;
479 bits_read += bit_count_to_read;
480 const result = huffman_tree.query(huffman_tree_index, prefix) catch |err| {
481 return err;
482 };
483
484 switch (result) {
485 .symbol => |sym| {
486 dest[i] = sym;
487 bit_count_to_read = starting_bit_count;
488 bits_read = 0;
489 huffman_tree_index = huffman_tree.symbol_count_minus_one;
490 break;
491 },
492 .index => |index| {
493 huffman_tree_index = index;
494 const bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
495 huffman_tree.nodes[index].weight,
496 max_bit_count,
497 );
498 bit_count_to_read = bit_count - bits_read;
499 },
500 }
501 }
502 }
503 self.literal_written_count += len;
504 },
505 }
506 }
507
508 /// Decode literals into `dest`; see `decodeLiteralsSlice()`.
509 pub fn decodeLiteralsRingBuffer(
510 self: *DecodeState,
511 dest: *RingBuffer,
512 len: usize,
513 ) DecodeLiteralsError!void {
514 if (self.literal_written_count + len > self.literal_header.regenerated_size)
515 return error.MalformedLiteralsLength;
516
517 switch (self.literal_header.block_type) {
518 .raw => {
519 const literals_end = self.literal_written_count + len;
520 const literal_data = self.literal_streams.one[self.literal_written_count..literals_end];
521 dest.writeSliceAssumeCapacity(literal_data);
522 self.literal_written_count += len;
523 },
524 .rle => {
525 var i: usize = 0;
526 while (i < len) : (i += 1) {
527 dest.writeAssumeCapacity(self.literal_streams.one[0]);
528 }
529 self.literal_written_count += len;
530 },
531 .compressed, .treeless => {
532 // const written_bytes_per_stream = (literals.header.regenerated_size + 3) / 4;
533 const huffman_tree = self.huffman_tree orelse unreachable;
534 const max_bit_count = huffman_tree.max_bit_count;
535 const starting_bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
536 huffman_tree.nodes[huffman_tree.symbol_count_minus_one].weight,
537 max_bit_count,
538 );
539 var bits_read: u4 = 0;
540 var huffman_tree_index: usize = huffman_tree.symbol_count_minus_one;
541 var bit_count_to_read: u4 = starting_bit_count;
542 var i: usize = 0;
543 while (i < len) : (i += 1) {
544 var prefix: u16 = 0;
545 while (true) {
546 const new_bits = try self.readLiteralsBits(u16, bit_count_to_read);
547 prefix <<= bit_count_to_read;
548 prefix |= new_bits;
549 bits_read += bit_count_to_read;
550 const result = try huffman_tree.query(huffman_tree_index, prefix);
551
552 switch (result) {
553 .symbol => |sym| {
554 dest.writeAssumeCapacity(sym);
555 bit_count_to_read = starting_bit_count;
556 bits_read = 0;
557 huffman_tree_index = huffman_tree.symbol_count_minus_one;
558 break;
559 },
560 .index => |index| {
561 huffman_tree_index = index;
562 const bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
563 huffman_tree.nodes[index].weight,
564 max_bit_count,
565 );
566 bit_count_to_read = bit_count - bits_read;
567 },
568 }
569 }
570 }
571 self.literal_written_count += len;
572 },
573 }
574 }
575
576 fn getCode(self: *DecodeState, comptime choice: DataType) u32 {
577 return switch (@field(self, @tagName(choice)).table) {
578 .rle => |value| value,
579 .fse => |table| table[@field(self, @tagName(choice)).state].symbol,
580 };
581 }
582};
583
58469pub fn computeChecksum(hasher: *std.hash.XxHash64) u32 {
58570 const hash = hasher.final();
58671 return @intCast(u32, hash & 0xFFFFFFFF);
......@@ -589,7 +74,7 @@ pub fn computeChecksum(hasher: *std.hash.XxHash64) u32 {
58974const FrameError = error{
59075 DictionaryIdFlagUnsupported,
59176 ChecksumFailure,
592} || InvalidBit || DecodeBlockError;
77} || InvalidBit || block.Error;
59378
59479/// Decode a Zstandard frame from `src` into `dest`, returning the number of
59580/// bytes read from `src` and written to `dest`; if the frame does not declare
......@@ -695,15 +180,15 @@ pub fn decodeZStandardFrameAlloc(
695180 var match_fse_data: [types.compressed_block.table_size_max.match]Table.Fse = undefined;
696181 var offset_fse_data: [types.compressed_block.table_size_max.offset]Table.Fse = undefined;
697182
698 var block_header = try decodeBlockHeaderSlice(src[consumed_count..]);
183 var block_header = try block.decodeBlockHeaderSlice(src[consumed_count..]);
699184 consumed_count += 3;
700 var decode_state = DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
185 var decode_state = block.DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
701186 while (true) : ({
702 block_header = try decodeBlockHeaderSlice(src[consumed_count..]);
187 block_header = try block.decodeBlockHeaderSlice(src[consumed_count..]);
703188 consumed_count += 3;
704189 }) {
705190 if (block_header.block_size > frame_context.block_size_max) return error.BlockSizeOverMaximum;
706 const written_size = try decodeBlockRingBuffer(
191 const written_size = try block.decodeBlockRingBuffer(
707192 &ring_buffer,
708193 src[consumed_count..],
709194 block_header,
......@@ -731,37 +216,28 @@ pub fn decodeZStandardFrameAlloc(
731216 return result.toOwnedSlice();
732217}
733218
734const DecodeBlockError = error{
735 BlockSizeOverMaximum,
736 MalformedBlockSize,
737 ReservedBlock,
738 MalformedRleBlock,
739 MalformedCompressedBlock,
740 EndOfStream,
741};
742
743219/// Convenience wrapper for decoding all blocks in a frame; see `decodeBlock()`.
744pub fn decodeFrameBlocks(
220fn decodeFrameBlocks(
745221 dest: []u8,
746222 src: []const u8,
747223 consumed_count: *usize,
748224 hash: ?*std.hash.XxHash64,
749) DecodeBlockError!usize {
225) block.Error!usize {
750226 // These tables take 7680 bytes
751227 var literal_fse_data: [types.compressed_block.table_size_max.literal]Table.Fse = undefined;
752228 var match_fse_data: [types.compressed_block.table_size_max.match]Table.Fse = undefined;
753229 var offset_fse_data: [types.compressed_block.table_size_max.offset]Table.Fse = undefined;
754230
755 var block_header = try decodeBlockHeaderSlice(src);
231 var block_header = try block.decodeBlockHeaderSlice(src);
756232 var bytes_read: usize = 3;
757233 defer consumed_count.* += bytes_read;
758 var decode_state = DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
234 var decode_state = block.DecodeState.init(&literal_fse_data, &match_fse_data, &offset_fse_data);
759235 var written_count: usize = 0;
760236 while (true) : ({
761 block_header = try decodeBlockHeaderSlice(src[bytes_read..]);
237 block_header = try block.decodeBlockHeaderSlice(src[bytes_read..]);
762238 bytes_read += 3;
763239 }) {
764 const written_size = try decodeBlock(
240 const written_size = try block.decodeBlock(
765241 dest,
766242 src[bytes_read..],
767243 block_header,
......@@ -776,255 +252,6 @@ pub fn decodeFrameBlocks(
776252 return written_count;
777253}
778254
779/// Decode a single block from `src` into `dest`. The beginning of `src` should
780/// be the start of the block content (i.e. directly after the block header).
781/// Increments `consumed_count` by the number of bytes read from `src` to decode
782/// the block and returns the decompressed size of the block.
783pub fn decodeBlock(
784 dest: []u8,
785 src: []const u8,
786 block_header: frame.ZStandard.Block.Header,
787 decode_state: *DecodeState,
788 consumed_count: *usize,
789 written_count: usize,
790) DecodeBlockError!usize {
791 const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB
792 const block_size = block_header.block_size;
793 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
794 switch (block_header.block_type) {
795 .raw => {
796 if (src.len < block_size) return error.MalformedBlockSize;
797 const data = src[0..block_size];
798 std.mem.copy(u8, dest[written_count..], data);
799 consumed_count.* += block_size;
800 return block_size;
801 },
802 .rle => {
803 if (src.len < 1) return error.MalformedRleBlock;
804 var write_pos: usize = written_count;
805 while (write_pos < block_size + written_count) : (write_pos += 1) {
806 dest[write_pos] = src[0];
807 }
808 consumed_count.* += 1;
809 return block_size;
810 },
811 .compressed => {
812 if (src.len < block_size) return error.MalformedBlockSize;
813 var bytes_read: usize = 0;
814 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
815 return error.MalformedCompressedBlock;
816 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
817 const fbs_reader = fbs.reader();
818 const sequences_header = decodeSequencesHeader(fbs_reader) catch
819 return error.MalformedCompressedBlock;
820
821 decode_state.prepare(fbs_reader, literals, sequences_header) catch
822 return error.MalformedCompressedBlock;
823
824 bytes_read += fbs.pos;
825
826 var bytes_written: usize = 0;
827 if (sequences_header.sequence_count > 0) {
828 const bit_stream_bytes = src[bytes_read..block_size];
829 var bit_stream: ReverseBitReader = undefined;
830 bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock;
831
832 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
833
834 var sequence_size_limit = block_size_max;
835 var i: usize = 0;
836 while (i < sequences_header.sequence_count) : (i += 1) {
837 const write_pos = written_count + bytes_written;
838 const decompressed_size = decode_state.decodeSequenceSlice(
839 dest,
840 write_pos,
841 &bit_stream,
842 sequence_size_limit,
843 i == sequences_header.sequence_count - 1,
844 ) catch return error.MalformedCompressedBlock;
845 bytes_written += decompressed_size;
846 sequence_size_limit -= decompressed_size;
847 }
848
849 bytes_read += bit_stream_bytes.len;
850 }
851 if (bytes_read != block_size) return error.MalformedCompressedBlock;
852
853 if (decode_state.literal_written_count < literals.header.regenerated_size) {
854 const len = literals.header.regenerated_size - decode_state.literal_written_count;
855 decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], len) catch
856 return error.MalformedCompressedBlock;
857 bytes_written += len;
858 }
859
860 consumed_count.* += bytes_read;
861 return bytes_written;
862 },
863 .reserved => return error.ReservedBlock,
864 }
865}
866
867/// Decode a single block from `src` into `dest`; see `decodeBlock()`. Returns
868/// the size of the decompressed block, which can be used with `dest.sliceLast()`
869/// to get the decompressed bytes.
870pub fn decodeBlockRingBuffer(
871 dest: *RingBuffer,
872 src: []const u8,
873 block_header: frame.ZStandard.Block.Header,
874 decode_state: *DecodeState,
875 consumed_count: *usize,
876 block_size_max: usize,
877) DecodeBlockError!usize {
878 const block_size = block_header.block_size;
879 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
880 switch (block_header.block_type) {
881 .raw => {
882 if (src.len < block_size) return error.MalformedBlockSize;
883 const data = src[0..block_size];
884 dest.writeSliceAssumeCapacity(data);
885 consumed_count.* += block_size;
886 return block_size;
887 },
888 .rle => {
889 if (src.len < 1) return error.MalformedRleBlock;
890 var write_pos: usize = 0;
891 while (write_pos < block_size) : (write_pos += 1) {
892 dest.writeAssumeCapacity(src[0]);
893 }
894 consumed_count.* += 1;
895 return block_size;
896 },
897 .compressed => {
898 if (src.len < block_size) return error.MalformedBlockSize;
899 var bytes_read: usize = 0;
900 const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch
901 return error.MalformedCompressedBlock;
902 var fbs = std.io.fixedBufferStream(src[bytes_read..]);
903 const fbs_reader = fbs.reader();
904 const sequences_header = decodeSequencesHeader(fbs_reader) catch
905 return error.MalformedCompressedBlock;
906
907 decode_state.prepare(fbs_reader, literals, sequences_header) catch
908 return error.MalformedCompressedBlock;
909
910 bytes_read += fbs.pos;
911
912 var bytes_written: usize = 0;
913 if (sequences_header.sequence_count > 0) {
914 const bit_stream_bytes = src[bytes_read..block_size];
915 var bit_stream: ReverseBitReader = undefined;
916 bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock;
917
918 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
919
920 var sequence_size_limit = block_size_max;
921 var i: usize = 0;
922 while (i < sequences_header.sequence_count) : (i += 1) {
923 const decompressed_size = decode_state.decodeSequenceRingBuffer(
924 dest,
925 &bit_stream,
926 sequence_size_limit,
927 i == sequences_header.sequence_count - 1,
928 ) catch return error.MalformedCompressedBlock;
929 bytes_written += decompressed_size;
930 sequence_size_limit -= decompressed_size;
931 }
932
933 bytes_read += bit_stream_bytes.len;
934 }
935 if (bytes_read != block_size) return error.MalformedCompressedBlock;
936
937 if (decode_state.literal_written_count < literals.header.regenerated_size) {
938 const len = literals.header.regenerated_size - decode_state.literal_written_count;
939 decode_state.decodeLiteralsRingBuffer(dest, len) catch
940 return error.MalformedCompressedBlock;
941 bytes_written += len;
942 }
943
944 consumed_count.* += bytes_read;
945 if (bytes_written > block_size_max) return error.BlockSizeOverMaximum;
946 return bytes_written;
947 },
948 .reserved => return error.ReservedBlock,
949 }
950}
951
952/// Decode a single block from `source` into `dest`. Literal and sequence data
953/// from the block is copied into `literals_buffer` and `sequence_buffer`, which
954/// must be large enough or `error.LiteralsBufferTooSmall` and
955/// `error.SequenceBufferTooSmall` are returned (the maximum block size is an
956/// upper bound for the size of both buffers). See `decodeBlock`
957/// and `decodeBlockRingBuffer` for function that can decode a block without
958/// these extra copies.
959pub fn decodeBlockReader(
960 dest: *RingBuffer,
961 source: anytype,
962 block_header: frame.ZStandard.Block.Header,
963 decode_state: *DecodeState,
964 block_size_max: usize,
965 literals_buffer: []u8,
966 sequence_buffer: []u8,
967) !void {
968 const block_size = block_header.block_size;
969 var block_reader_limited = std.io.limitedReader(source, block_size);
970 const block_reader = block_reader_limited.reader();
971 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
972 switch (block_header.block_type) {
973 .raw => {
974 const slice = dest.sliceAt(dest.write_index, block_size);
975 try source.readNoEof(slice.first);
976 try source.readNoEof(slice.second);
977 dest.write_index = dest.mask2(dest.write_index + block_size);
978 },
979 .rle => {
980 const byte = try source.readByte();
981 var i: usize = 0;
982 while (i < block_size) : (i += 1) {
983 dest.writeAssumeCapacity(byte);
984 }
985 },
986 .compressed => {
987 const literals = try decodeLiteralsSection(block_reader, literals_buffer);
988 const sequences_header = try decodeSequencesHeader(block_reader);
989
990 try decode_state.prepare(block_reader, literals, sequences_header);
991
992 if (sequences_header.sequence_count > 0) {
993 if (sequence_buffer.len < block_reader_limited.bytes_left)
994 return error.SequenceBufferTooSmall;
995
996 const size = try block_reader.readAll(sequence_buffer);
997 var bit_stream: ReverseBitReader = undefined;
998 try bit_stream.init(sequence_buffer[0..size]);
999
1000 decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock;
1001
1002 var sequence_size_limit = block_size_max;
1003 var i: usize = 0;
1004 while (i < sequences_header.sequence_count) : (i += 1) {
1005 const decompressed_size = decode_state.decodeSequenceRingBuffer(
1006 dest,
1007 &bit_stream,
1008 sequence_size_limit,
1009 i == sequences_header.sequence_count - 1,
1010 ) catch return error.MalformedCompressedBlock;
1011 sequence_size_limit -= decompressed_size;
1012 }
1013 }
1014
1015 if (decode_state.literal_written_count < literals.header.regenerated_size) {
1016 const len = literals.header.regenerated_size - decode_state.literal_written_count;
1017 decode_state.decodeLiteralsRingBuffer(dest, len) catch
1018 return error.MalformedCompressedBlock;
1019 }
1020
1021 decode_state.literal_written_count = 0;
1022 assert(block_reader.readByte() == error.EndOfStream);
1023 },
1024 .reserved => return error.ReservedBlock,
1025 }
1026}
1027
1028255/// Decode the header of a skippable frame.
1029256pub fn decodeSkippableHeader(src: *const [8]u8) frame.Skippable.Header {
1030257 const magic = readInt(u32, src[0..4]);
......@@ -1090,673 +317,6 @@ pub fn decodeZStandardHeader(source: anytype) (error{EndOfStream} || InvalidBit)
1090317 return header;
1091318}
1092319
1093/// Decode the header of a block.
1094pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header {
1095 const last_block = src[0] & 1 == 1;
1096 const block_type = @intToEnum(frame.ZStandard.Block.Type, (src[0] & 0b110) >> 1);
1097 const block_size = ((src[0] & 0b11111000) >> 3) + (@as(u21, src[1]) << 5) + (@as(u21, src[2]) << 13);
1098 return .{
1099 .last_block = last_block,
1100 .block_type = block_type,
1101 .block_size = block_size,
1102 };
1103}
1104
1105pub fn decodeBlockHeaderSlice(src: []const u8) error{EndOfStream}!frame.ZStandard.Block.Header {
1106 if (src.len < 3) return error.EndOfStream;
1107 return decodeBlockHeader(src[0..3]);
1108}
1109
1110/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
1111/// number of bytes the section uses.
1112///
1113/// Errors:
1114/// - returns `error.MalformedLiteralsHeader` if the header is invalid
1115/// - returns `error.MalformedLiteralsSection` if there are errors decoding
1116pub fn decodeLiteralsSectionSlice(
1117 src: []const u8,
1118 consumed_count: *usize,
1119) (error{ MalformedLiteralsHeader, MalformedLiteralsSection, EndOfStream } || DecodeHuffmanError)!LiteralsSection {
1120 var bytes_read: usize = 0;
1121 const header = header: {
1122 var fbs = std.io.fixedBufferStream(src);
1123 defer bytes_read = fbs.pos;
1124 break :header decodeLiteralsHeader(fbs.reader()) catch return error.MalformedLiteralsHeader;
1125 };
1126 switch (header.block_type) {
1127 .raw => {
1128 if (src.len < bytes_read + header.regenerated_size) return error.MalformedLiteralsSection;
1129 const stream = src[bytes_read .. bytes_read + header.regenerated_size];
1130 consumed_count.* += header.regenerated_size + bytes_read;
1131 return LiteralsSection{
1132 .header = header,
1133 .huffman_tree = null,
1134 .streams = .{ .one = stream },
1135 };
1136 },
1137 .rle => {
1138 if (src.len < bytes_read + 1) return error.MalformedLiteralsSection;
1139 const stream = src[bytes_read .. bytes_read + 1];
1140 consumed_count.* += 1 + bytes_read;
1141 return LiteralsSection{
1142 .header = header,
1143 .huffman_tree = null,
1144 .streams = .{ .one = stream },
1145 };
1146 },
1147 .compressed, .treeless => {
1148 const huffman_tree_start = bytes_read;
1149 const huffman_tree = if (header.block_type == .compressed)
1150 try decodeHuffmanTreeSlice(src[bytes_read..], &bytes_read)
1151 else
1152 null;
1153 const huffman_tree_size = bytes_read - huffman_tree_start;
1154 const total_streams_size = @as(usize, header.compressed_size.?) - huffman_tree_size;
1155
1156 if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection;
1157 const stream_data = src[bytes_read .. bytes_read + total_streams_size];
1158
1159 const streams = try decodeStreams(header.size_format, stream_data);
1160 consumed_count.* += bytes_read + total_streams_size;
1161 return LiteralsSection{
1162 .header = header,
1163 .huffman_tree = huffman_tree,
1164 .streams = streams,
1165 };
1166 },
1167 }
1168}
1169
1170/// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the
1171/// number of bytes the section uses.
1172///
1173/// Errors:
1174/// - returns `error.MalformedLiteralsHeader` if the header is invalid
1175/// - returns `error.MalformedLiteralsSection` if there are errors decoding
1176pub fn decodeLiteralsSection(
1177 source: anytype,
1178 buffer: []u8,
1179) !LiteralsSection {
1180 const header = try decodeLiteralsHeader(source);
1181 switch (header.block_type) {
1182 .raw => {
1183 try source.readNoEof(buffer[0..header.regenerated_size]);
1184 return LiteralsSection{
1185 .header = header,
1186 .huffman_tree = null,
1187 .streams = .{ .one = buffer },
1188 };
1189 },
1190 .rle => {
1191 buffer[0] = try source.readByte();
1192 return LiteralsSection{
1193 .header = header,
1194 .huffman_tree = null,
1195 .streams = .{ .one = buffer[0..1] },
1196 };
1197 },
1198 .compressed, .treeless => {
1199 var counting_reader = std.io.countingReader(source);
1200 const huffman_tree = if (header.block_type == .compressed)
1201 try decodeHuffmanTree(counting_reader.reader(), buffer)
1202 else
1203 null;
1204 const huffman_tree_size = counting_reader.bytes_read;
1205 const total_streams_size = @as(usize, header.compressed_size.?) - @intCast(usize, huffman_tree_size);
1206
1207 if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall;
1208 try source.readNoEof(buffer[0..total_streams_size]);
1209 const stream_data = buffer[0..total_streams_size];
1210
1211 const streams = try decodeStreams(header.size_format, stream_data);
1212 return LiteralsSection{
1213 .header = header,
1214 .huffman_tree = huffman_tree,
1215 .streams = streams,
1216 };
1217 },
1218 }
1219}
1220
1221fn decodeStreams(size_format: u2, stream_data: []const u8) !LiteralsSection.Streams {
1222 if (size_format == 0) {
1223 return .{ .one = stream_data };
1224 }
1225
1226 if (stream_data.len < 6) return error.MalformedLiteralsSection;
1227
1228 const stream_1_length = @as(usize, readInt(u16, stream_data[0..2]));
1229 const stream_2_length = @as(usize, readInt(u16, stream_data[2..4]));
1230 const stream_3_length = @as(usize, readInt(u16, stream_data[4..6]));
1231
1232 const stream_1_start = 6;
1233 const stream_2_start = stream_1_start + stream_1_length;
1234 const stream_3_start = stream_2_start + stream_2_length;
1235 const stream_4_start = stream_3_start + stream_3_length;
1236
1237 return .{ .four = .{
1238 stream_data[stream_1_start .. stream_1_start + stream_1_length],
1239 stream_data[stream_2_start .. stream_2_start + stream_2_length],
1240 stream_data[stream_3_start .. stream_3_start + stream_3_length],
1241 stream_data[stream_4_start..],
1242 } };
1243}
1244
1245const DecodeHuffmanError = error{
1246 MalformedHuffmanTree,
1247 MalformedFseTable,
1248 MalformedAccuracyLog,
1249};
1250
1251fn decodeFseHuffmanTree(source: anytype, compressed_size: usize, buffer: []u8, weights: *[256]u4) !usize {
1252 var stream = std.io.limitedReader(source, compressed_size);
1253 var bit_reader = bitReader(stream.reader());
1254
1255 var entries: [1 << 6]Table.Fse = undefined;
1256 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
1257 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
1258 error.EndOfStream => return error.MalformedFseTable,
1259 };
1260 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
1261
1262 const amount = try stream.reader().readAll(buffer);
1263 var huff_bits: ReverseBitReader = undefined;
1264 huff_bits.init(buffer[0..amount]) catch return error.MalformedHuffmanTree;
1265
1266 return assignWeights(&huff_bits, accuracy_log, &entries, weights);
1267}
1268
1269fn decodeFseHuffmanTreeSlice(src: []const u8, compressed_size: usize, weights: *[256]u4) !usize {
1270 if (src.len < compressed_size) return error.MalformedHuffmanTree;
1271 var stream = std.io.fixedBufferStream(src[0..compressed_size]);
1272 var counting_reader = std.io.countingReader(stream.reader());
1273 var bit_reader = bitReader(counting_reader.reader());
1274
1275 var entries: [1 << 6]Table.Fse = undefined;
1276 const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) {
1277 error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e,
1278 error.EndOfStream => return error.MalformedFseTable,
1279 };
1280 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
1281
1282 const start_index = std.math.cast(usize, counting_reader.bytes_read) orelse return error.MalformedHuffmanTree;
1283 var huff_data = src[start_index..compressed_size];
1284 var huff_bits: ReverseBitReader = undefined;
1285 huff_bits.init(huff_data) catch return error.MalformedHuffmanTree;
1286
1287 return assignWeights(&huff_bits, accuracy_log, &entries, weights);
1288}
1289
1290fn assignWeights(huff_bits: *ReverseBitReader, accuracy_log: usize, entries: *[1 << 6]Table.Fse, weights: *[256]u4) !usize {
1291 var i: usize = 0;
1292 var even_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
1293 var odd_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree;
1294
1295 while (i < 255) {
1296 const even_data = entries[even_state];
1297 var read_bits: usize = 0;
1298 const even_bits = huff_bits.readBits(u32, even_data.bits, &read_bits) catch unreachable;
1299 weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree;
1300 i += 1;
1301 if (read_bits < even_data.bits) {
1302 weights[i] = std.math.cast(u4, entries[odd_state].symbol) orelse return error.MalformedHuffmanTree;
1303 i += 1;
1304 break;
1305 }
1306 even_state = even_data.baseline + even_bits;
1307
1308 read_bits = 0;
1309 const odd_data = entries[odd_state];
1310 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;
1311 weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree;
1312 i += 1;
1313 if (read_bits < odd_data.bits) {
1314 if (i == 256) return error.MalformedHuffmanTree;
1315 weights[i] = std.math.cast(u4, entries[even_state].symbol) orelse return error.MalformedHuffmanTree;
1316 i += 1;
1317 break;
1318 }
1319 odd_state = odd_data.baseline + odd_bits;
1320 } else return error.MalformedHuffmanTree;
1321
1322 return i + 1; // stream contains all but the last symbol
1323}
1324
1325fn decodeDirectHuffmanTree(source: anytype, encoded_symbol_count: usize, weights: *[256]u4) !usize {
1326 const weights_byte_count = (encoded_symbol_count + 1) / 2;
1327 var i: usize = 0;
1328 while (i < weights_byte_count) : (i += 1) {
1329 const byte = try source.readByte();
1330 weights[2 * i] = @intCast(u4, byte >> 4);
1331 weights[2 * i + 1] = @intCast(u4, byte & 0xF);
1332 }
1333 return encoded_symbol_count + 1;
1334}
1335
1336fn assignSymbols(weight_sorted_prefixed_symbols: []LiteralsSection.HuffmanTree.PrefixedSymbol, weights: [256]u4) usize {
1337 for (weight_sorted_prefixed_symbols) |_, i| {
1338 weight_sorted_prefixed_symbols[i] = .{
1339 .symbol = @intCast(u8, i),
1340 .weight = undefined,
1341 .prefix = undefined,
1342 };
1343 }
1344
1345 std.sort.sort(
1346 LiteralsSection.HuffmanTree.PrefixedSymbol,
1347 weight_sorted_prefixed_symbols,
1348 weights,
1349 lessThanByWeight,
1350 );
1351
1352 var prefix: u16 = 0;
1353 var prefixed_symbol_count: usize = 0;
1354 var sorted_index: usize = 0;
1355 const symbol_count = weight_sorted_prefixed_symbols.len;
1356 while (sorted_index < symbol_count) {
1357 var symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
1358 const weight = weights[symbol];
1359 if (weight == 0) {
1360 sorted_index += 1;
1361 continue;
1362 }
1363
1364 while (sorted_index < symbol_count) : ({
1365 sorted_index += 1;
1366 prefixed_symbol_count += 1;
1367 prefix += 1;
1368 }) {
1369 symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
1370 if (weights[symbol] != weight) {
1371 prefix = ((prefix - 1) >> (weights[symbol] - weight)) + 1;
1372 break;
1373 }
1374 weight_sorted_prefixed_symbols[prefixed_symbol_count].symbol = symbol;
1375 weight_sorted_prefixed_symbols[prefixed_symbol_count].prefix = prefix;
1376 weight_sorted_prefixed_symbols[prefixed_symbol_count].weight = weight;
1377 }
1378 }
1379 return prefixed_symbol_count;
1380}
1381
1382fn buildHuffmanTree(weights: *[256]u4, symbol_count: usize) LiteralsSection.HuffmanTree {
1383 var weight_power_sum: u16 = 0;
1384 for (weights[0 .. symbol_count - 1]) |value| {
1385 if (value > 0) {
1386 weight_power_sum += @as(u16, 1) << (value - 1);
1387 }
1388 }
1389
1390 // advance to next power of two (even if weight_power_sum is a power of 2)
1391 const max_number_of_bits = std.math.log2_int(u16, weight_power_sum) + 1;
1392 const next_power_of_two = @as(u16, 1) << max_number_of_bits;
1393 weights[symbol_count - 1] = std.math.log2_int(u16, next_power_of_two - weight_power_sum) + 1;
1394
1395 var weight_sorted_prefixed_symbols: [256]LiteralsSection.HuffmanTree.PrefixedSymbol = undefined;
1396 const prefixed_symbol_count = assignSymbols(weight_sorted_prefixed_symbols[0..symbol_count], weights.*);
1397 const tree = LiteralsSection.HuffmanTree{
1398 .max_bit_count = max_number_of_bits,
1399 .symbol_count_minus_one = @intCast(u8, prefixed_symbol_count - 1),
1400 .nodes = weight_sorted_prefixed_symbols,
1401 };
1402 return tree;
1403}
1404
1405fn decodeHuffmanTree(source: anytype, buffer: []u8) !LiteralsSection.HuffmanTree {
1406 const header = try source.readByte();
1407 var weights: [256]u4 = undefined;
1408 const symbol_count = if (header < 128)
1409 // FSE compressed weights
1410 try decodeFseHuffmanTree(source, header, buffer, &weights)
1411 else
1412 try decodeDirectHuffmanTree(source, header - 127, &weights);
1413
1414 return buildHuffmanTree(&weights, symbol_count);
1415}
1416
1417fn decodeHuffmanTreeSlice(src: []const u8, consumed_count: *usize) (error{EndOfStream} || DecodeHuffmanError)!LiteralsSection.HuffmanTree {
1418 if (src.len == 0) return error.MalformedHuffmanTree;
1419 const header = src[0];
1420 var bytes_read: usize = 1;
1421 var weights: [256]u4 = undefined;
1422 const symbol_count = if (header < 128) count: {
1423 // FSE compressed weights
1424 bytes_read += header;
1425 break :count try decodeFseHuffmanTreeSlice(src[1..], header, &weights);
1426 } else count: {
1427 var fbs = std.io.fixedBufferStream(src[1..]);
1428 defer bytes_read += fbs.pos;
1429 break :count try decodeDirectHuffmanTree(fbs.reader(), header - 127, &weights);
1430 };
1431
1432 consumed_count.* += bytes_read;
1433 return buildHuffmanTree(&weights, symbol_count);
1434}
1435
1436fn lessThanByWeight(
1437 weights: [256]u4,
1438 lhs: LiteralsSection.HuffmanTree.PrefixedSymbol,
1439 rhs: LiteralsSection.HuffmanTree.PrefixedSymbol,
1440) bool {
1441 // NOTE: this function relies on the use of a stable sorting algorithm,
1442 // otherwise a special case of if (weights[lhs] == weights[rhs]) return lhs < rhs;
1443 // should be added
1444 return weights[lhs.symbol] < weights[rhs.symbol];
1445}
1446
1447/// Decode a literals section header.
1448pub fn decodeLiteralsHeader(source: anytype) !LiteralsSection.Header {
1449 const byte0 = try source.readByte();
1450 const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11);
1451 const size_format = @intCast(u2, (byte0 & 0b1100) >> 2);
1452 var regenerated_size: u20 = undefined;
1453 var compressed_size: ?u18 = null;
1454 switch (block_type) {
1455 .raw, .rle => {
1456 switch (size_format) {
1457 0, 2 => {
1458 regenerated_size = byte0 >> 3;
1459 },
1460 1 => regenerated_size = (byte0 >> 4) + (@as(u20, try source.readByte()) << 4),
1461 3 => regenerated_size = (byte0 >> 4) +
1462 (@as(u20, try source.readByte()) << 4) +
1463 (@as(u20, try source.readByte()) << 12),
1464 }
1465 },
1466 .compressed, .treeless => {
1467 const byte1 = try source.readByte();
1468 const byte2 = try source.readByte();
1469 switch (size_format) {
1470 0, 1 => {
1471 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);
1472 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);
1473 },
1474 2 => {
1475 const byte3 = try source.readByte();
1476 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);
1477 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);
1478 },
1479 3 => {
1480 const byte3 = try source.readByte();
1481 const byte4 = try source.readByte();
1482 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);
1483 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);
1484 },
1485 }
1486 },
1487 }
1488 return LiteralsSection.Header{
1489 .block_type = block_type,
1490 .size_format = size_format,
1491 .regenerated_size = regenerated_size,
1492 .compressed_size = compressed_size,
1493 };
1494}
1495
1496/// Decode a sequences section header.
1497///
1498/// Errors:
1499/// - returns `error.ReservedBitSet` is the reserved bit is set
1500/// - returns `error.MalformedSequencesHeader` if the header is invalid
1501pub fn decodeSequencesHeader(
1502 source: anytype,
1503) !SequencesSection.Header {
1504 var sequence_count: u24 = undefined;
1505
1506 const byte0 = try source.readByte();
1507 if (byte0 == 0) {
1508 return SequencesSection.Header{
1509 .sequence_count = 0,
1510 .offsets = undefined,
1511 .match_lengths = undefined,
1512 .literal_lengths = undefined,
1513 };
1514 } else if (byte0 < 128) {
1515 sequence_count = byte0;
1516 } else if (byte0 < 255) {
1517 sequence_count = (@as(u24, (byte0 - 128)) << 8) + try source.readByte();
1518 } else {
1519 sequence_count = (try source.readByte()) + (@as(u24, try source.readByte()) << 8) + 0x7F00;
1520 }
1521
1522 const compression_modes = try source.readByte();
1523
1524 const matches_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00001100) >> 2);
1525 const offsets_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b00110000) >> 4);
1526 const literal_mode = @intToEnum(SequencesSection.Header.Mode, (compression_modes & 0b11000000) >> 6);
1527 if (compression_modes & 0b11 != 0) return error.ReservedBitSet;
1528
1529 return SequencesSection.Header{
1530 .sequence_count = sequence_count,
1531 .offsets = offsets_mode,
1532 .match_lengths = matches_mode,
1533 .literal_lengths = literal_mode,
1534 };
1535}
1536
1537fn buildFseTable(values: []const u16, entries: []Table.Fse) !void {
1538 const total_probability = @intCast(u16, entries.len);
1539 const accuracy_log = std.math.log2_int(u16, total_probability);
1540 assert(total_probability <= 1 << 9);
1541
1542 var less_than_one_count: usize = 0;
1543 for (values) |value, i| {
1544 if (value == 0) {
1545 entries[entries.len - 1 - less_than_one_count] = Table.Fse{
1546 .symbol = @intCast(u8, i),
1547 .baseline = 0,
1548 .bits = accuracy_log,
1549 };
1550 less_than_one_count += 1;
1551 }
1552 }
1553
1554 var position: usize = 0;
1555 var temp_states: [1 << 9]u16 = undefined;
1556 for (values) |value, symbol| {
1557 if (value == 0 or value == 1) continue;
1558 const probability = value - 1;
1559
1560 const state_share_dividend = std.math.ceilPowerOfTwo(u16, probability) catch
1561 return error.MalformedFseTable;
1562 const share_size = @divExact(total_probability, state_share_dividend);
1563 const double_state_count = state_share_dividend - probability;
1564 const single_state_count = probability - double_state_count;
1565 const share_size_log = std.math.log2_int(u16, share_size);
1566
1567 var i: u16 = 0;
1568 while (i < probability) : (i += 1) {
1569 temp_states[i] = @intCast(u16, position);
1570 position += (entries.len >> 1) + (entries.len >> 3) + 3;
1571 position &= entries.len - 1;
1572 while (position >= entries.len - less_than_one_count) {
1573 position += (entries.len >> 1) + (entries.len >> 3) + 3;
1574 position &= entries.len - 1;
1575 }
1576 }
1577 std.sort.sort(u16, temp_states[0..probability], {}, std.sort.asc(u16));
1578 i = 0;
1579 while (i < probability) : (i += 1) {
1580 entries[temp_states[i]] = if (i < double_state_count) Table.Fse{
1581 .symbol = @intCast(u8, symbol),
1582 .bits = share_size_log + 1,
1583 .baseline = single_state_count * share_size + i * 2 * share_size,
1584 } else Table.Fse{
1585 .symbol = @intCast(u8, symbol),
1586 .bits = share_size_log,
1587 .baseline = (i - double_state_count) * share_size,
1588 };
1589 }
1590 }
1591}
1592
1593fn decodeFseTable(
1594 bit_reader: anytype,
1595 expected_symbol_count: usize,
1596 max_accuracy_log: u4,
1597 entries: []Table.Fse,
1598) !usize {
1599 const accuracy_log_biased = try bit_reader.readBitsNoEof(u4, 4);
1600 if (accuracy_log_biased > max_accuracy_log -| 5) return error.MalformedAccuracyLog;
1601 const accuracy_log = accuracy_log_biased + 5;
1602
1603 var values: [256]u16 = undefined;
1604 var value_count: usize = 0;
1605
1606 const total_probability = @as(u16, 1) << accuracy_log;
1607 var accumulated_probability: u16 = 0;
1608
1609 while (accumulated_probability < total_probability) {
1610 // WARNING: The RFC in poorly worded, and would suggest std.math.log2_int_ceil is correct here,
1611 // but power of two (remaining probabilities + 1) need max bits set to 1 more.
1612 const max_bits = std.math.log2_int(u16, total_probability - accumulated_probability + 1) + 1;
1613 const small = try bit_reader.readBitsNoEof(u16, max_bits - 1);
1614
1615 const cutoff = (@as(u16, 1) << max_bits) - 1 - (total_probability - accumulated_probability + 1);
1616
1617 const value = if (small < cutoff)
1618 small
1619 else value: {
1620 const value_read = small + (try bit_reader.readBitsNoEof(u16, 1) << (max_bits - 1));
1621 break :value if (value_read < @as(u16, 1) << (max_bits - 1))
1622 value_read
1623 else
1624 value_read - cutoff;
1625 };
1626
1627 accumulated_probability += if (value != 0) value - 1 else 1;
1628
1629 values[value_count] = value;
1630 value_count += 1;
1631
1632 if (value == 1) {
1633 while (true) {
1634 const repeat_flag = try bit_reader.readBitsNoEof(u2, 2);
1635 var i: usize = 0;
1636 while (i < repeat_flag) : (i += 1) {
1637 values[value_count] = 1;
1638 value_count += 1;
1639 }
1640 if (repeat_flag < 3) break;
1641 }
1642 }
1643 }
1644 bit_reader.alignToByte();
1645
1646 if (value_count < 2) return error.MalformedFseTable;
1647 if (accumulated_probability != total_probability) return error.MalformedFseTable;
1648 if (value_count > expected_symbol_count) return error.MalformedFseTable;
1649
1650 const table_size = total_probability;
1651
1652 try buildFseTable(values[0..value_count], entries[0..table_size]);
1653 return table_size;
1654}
1655
1656const ReversedByteReader = struct {
1657 remaining_bytes: usize,
1658 bytes: []const u8,
1659
1660 const Reader = std.io.Reader(*ReversedByteReader, error{}, readFn);
1661
1662 fn init(bytes: []const u8) ReversedByteReader {
1663 return .{
1664 .bytes = bytes,
1665 .remaining_bytes = bytes.len,
1666 };
1667 }
1668
1669 fn reader(self: *ReversedByteReader) Reader {
1670 return .{ .context = self };
1671 }
1672
1673 fn readFn(ctx: *ReversedByteReader, buffer: []u8) !usize {
1674 if (ctx.remaining_bytes == 0) return 0;
1675 const byte_index = ctx.remaining_bytes - 1;
1676 buffer[0] = ctx.bytes[byte_index];
1677 // buffer[0] = @bitReverse(ctx.bytes[byte_index]);
1678 ctx.remaining_bytes = byte_index;
1679 return 1;
1680 }
1681};
1682
1683/// A bit reader for reading the reversed bit streams used to encode
1684/// FSE compressed data.
1685pub const ReverseBitReader = struct {
1686 byte_reader: ReversedByteReader,
1687 bit_reader: std.io.BitReader(.Big, ReversedByteReader.Reader),
1688
1689 pub fn init(self: *ReverseBitReader, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
1690 self.byte_reader = ReversedByteReader.init(bytes);
1691 self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader());
1692 while (0 == self.readBitsNoEof(u1, 1) catch return error.BitStreamHasNoStartBit) {}
1693 }
1694
1695 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) error{EndOfStream}!U {
1696 return self.bit_reader.readBitsNoEof(U, num_bits);
1697 }
1698
1699 pub fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) error{}!U {
1700 return try self.bit_reader.readBits(U, num_bits, out_bits);
1701 }
1702
1703 pub fn alignToByte(self: *@This()) void {
1704 self.bit_reader.alignToByte();
1705 }
1706};
1707
1708fn BitReader(comptime Reader: type) type {
1709 return struct {
1710 underlying: std.io.BitReader(.Little, Reader),
1711
1712 fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) !U {
1713 return self.underlying.readBitsNoEof(U, num_bits);
1714 }
1715
1716 fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) !U {
1717 return self.underlying.readBits(U, num_bits, out_bits);
1718 }
1719
1720 fn alignToByte(self: *@This()) void {
1721 self.underlying.alignToByte();
1722 }
1723 };
1724}
1725
1726pub fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) {
1727 return .{ .underlying = std.io.bitReader(.Little, reader) };
1728}
1729
1730320test {
1731321 std.testing.refAllDecls(@This());
1732322}
1733
1734test buildFseTable {
1735 const literals_length_default_values = [36]u16{
1736 5, 4, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2,
1737 3, 3, 3, 3, 3, 3, 3, 3, 3, 4, 3, 2, 2, 2, 2, 2,
1738 0, 0, 0, 0,
1739 };
1740
1741 const match_lengths_default_values = [53]u16{
1742 2, 5, 4, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2,
1743 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
1744 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 0, 0,
1745 0, 0, 0, 0, 0,
1746 };
1747
1748 const offset_codes_default_values = [29]u16{
1749 2, 2, 2, 2, 2, 2, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2,
1750 2, 2, 2, 2, 2, 2, 2, 2, 0, 0, 0, 0, 0,
1751 };
1752
1753 var entries: [64]Table.Fse = undefined;
1754 try buildFseTable(&literals_length_default_values, &entries);
1755 try std.testing.expectEqualSlices(Table.Fse, types.compressed_block.predefined_literal_fse_table.fse, &entries);
1756
1757 try buildFseTable(&match_lengths_default_values, &entries);
1758 try std.testing.expectEqualSlices(Table.Fse, types.compressed_block.predefined_match_fse_table.fse, &entries);
1759
1760 try buildFseTable(&offset_codes_default_values, entries[0..32]);
1761 try std.testing.expectEqualSlices(Table.Fse, types.compressed_block.predefined_offset_fse_table.fse, entries[0..32]);
1762}
lib/std/compress/zstandard/readers.zig created+75
......@@ -0,0 +1,75 @@
1const std = @import("std");
2
3pub const ReversedByteReader = struct {
4 remaining_bytes: usize,
5 bytes: []const u8,
6
7 const Reader = std.io.Reader(*ReversedByteReader, error{}, readFn);
8
9 pub fn init(bytes: []const u8) ReversedByteReader {
10 return .{
11 .bytes = bytes,
12 .remaining_bytes = bytes.len,
13 };
14 }
15
16 pub fn reader(self: *ReversedByteReader) Reader {
17 return .{ .context = self };
18 }
19
20 fn readFn(ctx: *ReversedByteReader, buffer: []u8) !usize {
21 if (ctx.remaining_bytes == 0) return 0;
22 const byte_index = ctx.remaining_bytes - 1;
23 buffer[0] = ctx.bytes[byte_index];
24 // buffer[0] = @bitReverse(ctx.bytes[byte_index]);
25 ctx.remaining_bytes = byte_index;
26 return 1;
27 }
28};
29
30/// A bit reader for reading the reversed bit streams used to encode
31/// FSE compressed data.
32pub const ReverseBitReader = struct {
33 byte_reader: ReversedByteReader,
34 bit_reader: std.io.BitReader(.Big, ReversedByteReader.Reader),
35
36 pub fn init(self: *ReverseBitReader, bytes: []const u8) error{BitStreamHasNoStartBit}!void {
37 self.byte_reader = ReversedByteReader.init(bytes);
38 self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader());
39 while (0 == self.readBitsNoEof(u1, 1) catch return error.BitStreamHasNoStartBit) {}
40 }
41
42 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) error{EndOfStream}!U {
43 return self.bit_reader.readBitsNoEof(U, num_bits);
44 }
45
46 pub fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) error{}!U {
47 return try self.bit_reader.readBits(U, num_bits, out_bits);
48 }
49
50 pub fn alignToByte(self: *@This()) void {
51 self.bit_reader.alignToByte();
52 }
53};
54
55pub fn BitReader(comptime Reader: type) type {
56 return struct {
57 underlying: std.io.BitReader(.Little, Reader),
58
59 pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) !U {
60 return self.underlying.readBitsNoEof(U, num_bits);
61 }
62
63 pub fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) !U {
64 return self.underlying.readBits(U, num_bits, out_bits);
65 }
66
67 pub fn alignToByte(self: *@This()) void {
68 self.underlying.alignToByte();
69 }
70 };
71}
72
73pub fn bitReader(reader: anytype) BitReader(@TypeOf(reader)) {
74 return .{ .underlying = std.io.bitReader(.Little, reader) };
75}