1const Decompress = @This();
2const std = @import("std");
3const assert = std.debug.assert;
4const Reader = std.Io.Reader;
5const Limit = std.Io.Limit;
6const zstd = @import("../zstd.zig");
7const Writer = std.Io.Writer;
8
9input: *Reader,
10reader: Reader,
11state: State,
12verify_checksum: bool,
13window_len: u32,
14err: ?Error = null,
15
16const State = union(enum) {
17 new_frame,
18 in_frame: InFrame,
19 skipping_frame: usize,
20
21 const InFrame = struct {
22 frame: Frame,
23 checksum: ?u32,
24 decompressed_size: usize,
25 decode: Frame.Zstandard.Decode,
26 };
27};
28
29pub const Options = struct {
30 /// Verifying checksums is not implemented yet and will cause a panic if
31 /// you set this to true.
32 verify_checksum: bool = false,
33
34 /// The output buffer is asserted to have capacity for `window_len` plus
35 /// `zstd.block_size_max`.
36 ///
37 /// If `window_len` is too small, then some streams will fail to decompress
38 /// with `error.OutputBufferUndersize`.
39 window_len: u32 = zstd.default_window_len,
40};
41
42pub const Error = error{
43 BadMagic,
44 BlockOversize,
45 ChecksumFailure,
46 ContentOversize,
47 DictionaryIdFlagUnsupported,
48 EndOfStream,
49 HuffmanTreeIncomplete,
50 InvalidBitStream,
51 MalformedAccuracyLog,
52 MalformedBlock,
53 MalformedCompressedBlock,
54 MalformedFrame,
55 MalformedFseBits,
56 MalformedFseTable,
57 MalformedHuffmanTree,
58 MalformedLiteralsHeader,
59 MalformedLiteralsLength,
60 MalformedLiteralsSection,
61 MalformedSequence,
62 MissingStartBit,
63 OutputBufferUndersize,
64 InputBufferUndersize,
65 ReadFailed,
66 RepeatModeFirst,
67 ReservedBitSet,
68 ReservedBlock,
69 SequenceBufferUndersize,
70 TreelessLiteralsFirst,
71 UnexpectedEndOfLiteralStream,
72 WindowOversize,
73 WindowSizeUnknown,
74};
75
76const direct_vtable: Reader.VTable = .{
77 .stream = streamDirect,
78 .rebase = rebaseFallible,
79 .discard = discardDirect,
80 .readVec = readVec,
81};
82
83const indirect_vtable: Reader.VTable = .{
84 .stream = streamIndirect,
85 .rebase = rebaseFallible,
86 .discard = discardIndirect,
87 .readVec = readVec,
88};
89
90/// When connecting `reader` to a `Writer`, `buffer` should be empty, and
91/// `Writer.buffer` capacity has requirements based on `Options.window_len`.
92///
93/// Otherwise, `buffer` has those requirements.
94pub fn init(input: *Reader, buffer: []u8, options: Options) Decompress {
95 if (buffer.len != 0) assert(buffer.len >= options.window_len + zstd.block_size_max);
96 return .{
97 .input = input,
98 .state = .new_frame,
99 .verify_checksum = options.verify_checksum,
100 .window_len = options.window_len,
101 .reader = .{
102 .vtable = if (buffer.len == 0) &direct_vtable else &indirect_vtable,
103 .buffer = buffer,
104 .seek = 0,
105 .end = 0,
106 },
107 };
108}
109
110fn streamDirect(r: *Reader, w: *Writer, limit: std.Io.Limit) Reader.StreamError!usize {
111 const d: *Decompress = @alignCast(@fieldParentPtr("reader", r));
112 return stream(d, w, limit);
113}
114
115fn streamIndirect(r: *Reader, w: *Writer, limit: std.Io.Limit) Reader.StreamError!usize {
116 const d: *Decompress = @alignCast(@fieldParentPtr("reader", r));
117 _ = limit;
118 _ = w;
119 return streamIndirectInner(d);
120}
121
122fn rebaseFallible(r: *Reader, capacity: usize) Reader.RebaseError!void {
123 rebase(r, capacity);
124}
125
126// Rebase the buffer, keeping at least the sliding window (`d.window_len` bytes) buffered
127fn rebase(r: *Reader, capacity: usize) void {
128 const d: *Decompress = @alignCast(@fieldParentPtr("reader", r));
129 // `capacity` must fit in the buffer along with the required sliding window
130 assert(capacity <= r.buffer.len - d.window_len);
131 // According to the vtable contract, this function will only be called if the free space in the
132 // buffer cannot already fit `capacity` bytes
133 assert(r.end + capacity > r.buffer.len);
134 const discard_n = @min(r.seek, r.end - d.window_len);
135 const keep = r.buffer[discard_n..r.end];
136 @memmove(r.buffer[0..keep.len], keep);
137 r.end = keep.len;
138 r.seek -= discard_n;
139}
140
141/// Rebase `d.reader.buffer` as much as needed for a discard limited by `limit`
142fn rebaseForDiscard(d: *Decompress, limit: std.Io.Limit) void {
143 // Number of bytes desired to rebase, always rebase for at least block_size
144 const desire_n = limit.max(Limit.limited(zstd.block_size_max));
145 // Maximum number of bytes possible to rebase
146 const max_n = d.reader.buffer.len -| d.window_len;
147 // Number of bytes to rebase
148 const n = desire_n.minInt(max_n);
149
150 // Current buffer free space
151 const current_cap = d.reader.buffer.len - d.reader.end;
152 if (current_cap < n) {
153 rebase(&d.reader, n);
154 }
155}
156
157/// This could be improved so that when an amount is discarded that includes an
158/// entire frame, skip decoding that frame.
159fn discardDirect(r: *Reader, limit: std.Io.Limit) Reader.Error!usize {
160 const d: *Decompress = @alignCast(@fieldParentPtr("reader", r));
161 rebaseForDiscard(d, limit);
162 var writer: Writer = .{
163 .vtable = &.{
164 .drain = std.Io.Writer.Discarding.drain,
165 .sendFile = std.Io.Writer.Discarding.sendFile,
166 },
167 .buffer = r.buffer,
168 .end = r.end,
169 };
170 defer {
171 r.end = writer.end;
172 r.seek = r.end;
173 }
174 const n = r.stream(&writer, limit) catch |err| switch (err) {
175 error.WriteFailed => unreachable,
176 error.ReadFailed, error.EndOfStream => |e| return e,
177 };
178 assert(n <= @backingInt(limit));
179 return n;
180}
181
182fn discardIndirect(r: *Reader, limit: std.Io.Limit) Reader.Error!usize {
183 const d: *Decompress = @alignCast(@fieldParentPtr("reader", r));
184 rebaseForDiscard(d, limit);
185 var writer: Writer = .{
186 .buffer = r.buffer,
187 .end = r.end,
188 .vtable = &.{ .drain = Writer.unreachableDrain },
189 };
190 {
191 defer r.end = writer.end;
192 _ = stream(d, &writer, .limited(writer.buffer.len - writer.end)) catch |err| switch (err) {
193 error.WriteFailed => unreachable,
194 else => |e| return e,
195 };
196 }
197 const n = limit.minInt(r.end - r.seek);
198 r.seek += n;
199 return n;
200}
201
202fn readVec(r: *Reader, data: [][]u8) Reader.Error!usize {
203 _ = data;
204 const d: *Decompress = @alignCast(@fieldParentPtr("reader", r));
205 return streamIndirectInner(d);
206}
207
208fn streamIndirectInner(d: *Decompress) Reader.Error!usize {
209 const r = &d.reader;
210 if (r.buffer.len - r.end < zstd.block_size_max) rebase(r, zstd.block_size_max);
211 assert(r.buffer.len - r.end >= zstd.block_size_max);
212 var writer: Writer = .{
213 .buffer = r.buffer,
214 .end = r.end,
215 .vtable = &.{
216 .drain = Writer.unreachableDrain,
217 .rebase = Writer.unreachableRebase,
218 },
219 };
220 defer r.end = writer.end;
221 _ = stream(d, &writer, .limited(writer.buffer.len - writer.end)) catch |err| switch (err) {
222 error.WriteFailed => unreachable,
223 else => |e| return e,
224 };
225 return 0;
226}
227
228fn stream(d: *Decompress, w: *Writer, limit: Limit) Reader.StreamError!usize {
229 const in = d.input;
230
231 state: switch (d.state) {
232 .new_frame => {
233 // Only return EndOfStream when there are exactly 0 bytes remaining on the
234 // frame magic. Any partial magic bytes should be considered a failure.
235 in.fill(@sizeOf(Frame.Magic)) catch |err| switch (err) {
236 error.EndOfStream => {
237 if (in.bufferedLen() != 0) {
238 d.err = error.BadMagic;
239 return error.ReadFailed;
240 }
241 return err;
242 },
243 else => |e| return e,
244 };
245 const magic = try in.takeEnumNonexhaustive(Frame.Magic, .little);
246 initFrame(d, magic) catch |err| {
247 d.err = err;
248 return error.ReadFailed;
249 };
250 continue :state d.state;
251 },
252 .in_frame => |*in_frame| {
253 return readInFrame(d, w, limit, in_frame) catch |err| switch (err) {
254 error.ReadFailed, error.WriteFailed => |e| return e,
255 else => |e| {
256 d.err = e;
257 return error.ReadFailed;
258 },
259 };
260 },
261 .skipping_frame => |*remaining| {
262 const n = in.discard(.limited(remaining.*)) catch |err| {
263 d.err = err;
264 return error.ReadFailed;
265 };
266 remaining.* -= n;
267 if (remaining.* == 0) d.state = .new_frame;
268 return 0;
269 },
270 }
271}
272
273fn initFrame(d: *Decompress, magic: Frame.Magic) !void {
274 const in = d.input;
275 switch (magic.kind() orelse return error.BadMagic) {
276 .zstandard => {
277 const header = try Frame.Zstandard.Header.decode(in);
278 d.state = .{ .in_frame = .{
279 .frame = try Frame.init(header, d.window_len, d.verify_checksum),
280 .checksum = null,
281 .decompressed_size = 0,
282 .decode = .init,
283 } };
284 },
285 .skippable => {
286 const frame_size = try in.takeInt(u32, .little);
287 d.state = .{ .skipping_frame = frame_size };
288 },
289 }
290}
291
292fn readInFrame(d: *Decompress, w: *Writer, limit: Limit, state: *State.InFrame) !usize {
293 const in = d.input;
294 const window_len = d.window_len;
295
296 const block_header = try in.takeStruct(Frame.Zstandard.Block.Header, .little);
297 const block_size = block_header.size;
298 const frame_block_size_max = state.frame.block_size_max;
299 if (frame_block_size_max < block_size) return error.BlockOversize;
300 if (@backingInt(limit) < block_size) return error.OutputBufferUndersize;
301 var bytes_written: usize = 0;
302 switch (block_header.type) {
303 .raw => {
304 try in.streamExactPreserve(w, window_len, block_size);
305 bytes_written = block_size;
306 },
307 .rle => {
308 const byte = try in.takeByte();
309 try w.splatBytePreserve(window_len, byte, block_size);
310 bytes_written = block_size;
311 },
312 .compressed => {
313 var literals_buffer: [zstd.block_size_max]u8 = undefined;
314 var sequence_buffer: [zstd.block_size_max]u8 = undefined;
315 var remaining: Limit = .limited(block_size);
316 const literals = try LiteralsSection.decode(in, &remaining, &literals_buffer);
317 const sequences_header = try SequencesSection.Header.decode(in, &remaining);
318
319 const decode = &state.decode;
320 try decode.prepare(in, &remaining, literals, sequences_header);
321
322 {
323 if (sequence_buffer.len < @backingInt(remaining))
324 return error.SequenceBufferUndersize;
325 const seq_slice = remaining.slice(&sequence_buffer);
326 try in.readSliceAll(seq_slice);
327 var bit_stream = try ReverseBitReader.init(seq_slice);
328
329 if (sequences_header.sequence_count > 0) {
330 try decode.readInitialFseState(&bit_stream);
331
332 // Ensures the following calls to `decodeSequence` will not flush.
333 const dest = (try w.writableSliceGreedyPreserve(window_len, frame_block_size_max))[0..frame_block_size_max];
334 const write_pos = dest.ptr - w.buffer.ptr;
335 for (0..sequences_header.sequence_count - 1) |_| {
336 bytes_written += try decode.decodeSequence(w.buffer, write_pos + bytes_written, &bit_stream);
337 try decode.updateState(.literal, &bit_stream);
338 try decode.updateState(.match, &bit_stream);
339 try decode.updateState(.offset, &bit_stream);
340 }
341 bytes_written += try decode.decodeSequence(w.buffer, write_pos + bytes_written, &bit_stream);
342 if (bytes_written > dest.len) return error.MalformedSequence;
343 w.advance(bytes_written);
344 }
345
346 if (!bit_stream.isEmpty()) {
347 return error.MalformedCompressedBlock;
348 }
349 }
350
351 if (decode.literal_written_count < literals.header.regenerated_size) {
352 const len = literals.header.regenerated_size - decode.literal_written_count;
353 try decode.decodeLiterals(w, len);
354 decode.literal_written_count += len;
355 bytes_written += len;
356 }
357
358 switch (decode.literal_header.block_type) {
359 .treeless, .compressed => {
360 if (!decode.isLiteralStreamEmpty()) return error.MalformedCompressedBlock;
361 },
362 .raw, .rle => {},
363 }
364
365 if (bytes_written > frame_block_size_max) return error.BlockOversize;
366 },
367 .reserved => return error.ReservedBlock,
368 }
369
370 if (state.frame.hasher_opt) |*hasher| {
371 if (bytes_written > 0) {
372 _ = hasher;
373 @panic("TODO all those bytes written needed to go through the hasher too");
374 }
375 }
376
377 state.decompressed_size += bytes_written;
378
379 if (block_header.last) {
380 if (state.frame.has_checksum) {
381 const expected_checksum = try in.takeInt(u32, .little);
382 if (state.frame.hasher_opt) |*hasher| {
383 const actual_checksum: u32 = @truncate(hasher.final());
384 if (expected_checksum != actual_checksum) return error.ChecksumFailure;
385 }
386 }
387 if (state.frame.content_size) |content_size| {
388 if (content_size != state.decompressed_size) {
389 return error.MalformedFrame;
390 }
391 }
392 d.state = .new_frame;
393 } else if (state.frame.content_size) |content_size| {
394 if (state.decompressed_size > content_size) return error.MalformedFrame;
395 }
396
397 return bytes_written;
398}
399
400pub const Frame = struct {
401 hasher_opt: ?std.hash.XxHash64,
402 window_size: usize,
403 has_checksum: bool,
404 block_size_max: usize,
405 content_size: ?usize,
406
407 pub const Magic = enum(u32) {
408 zstandard = 0xFD2FB528,
409 _,
410
411 pub fn kind(m: Magic) ?Kind {
412 return switch (@backingInt(m)) {
413 @backingInt(Magic.zstandard) => .zstandard,
414 @backingInt(Skippable.magic_min)...@backingInt(Skippable.magic_max) => .skippable,
415 else => null,
416 };
417 }
418
419 pub fn isSkippable(m: Magic) bool {
420 return switch (@backingInt(m)) {
421 @backingInt(Skippable.magic_min)...@backingInt(Skippable.magic_max) => true,
422 else => false,
423 };
424 }
425 };
426
427 pub const Kind = enum { zstandard, skippable };
428
429 pub const Zstandard = struct {
430 pub const magic: Magic = .zstandard;
431
432 header: Header,
433 data_blocks: []Block,
434 checksum: ?u32,
435
436 pub const Header = struct {
437 descriptor: Descriptor,
438 window_descriptor: ?u8,
439 dictionary_id: ?u32,
440 content_size: ?u64,
441
442 pub const Descriptor = packed struct {
443 dictionary_id_flag: u2,
444 content_checksum_flag: bool,
445 reserved: bool,
446 unused: bool,
447 single_segment_flag: bool,
448 content_size_flag: u2,
449 };
450
451 pub const DecodeError = Reader.Error || error{ReservedBitSet};
452
453 pub fn decode(in: *Reader) DecodeError!Header {
454 const descriptor: Descriptor = @bitCast(try in.takeByte());
455
456 if (descriptor.reserved) return error.ReservedBitSet;
457
458 const window_descriptor: ?u8 = if (descriptor.single_segment_flag) null else try in.takeByte();
459
460 const dictionary_id: ?u32 = if (descriptor.dictionary_id_flag > 0) d: {
461 // if flag is 3 then field_size = 4, else field_size = flag
462 const field_size = (@as(u4, 1) << descriptor.dictionary_id_flag) >> 1;
463 break :d try in.takeVarInt(u32, .little, field_size);
464 } else null;
465
466 const content_size: ?u64 = if (descriptor.single_segment_flag or descriptor.content_size_flag > 0) c: {
467 const field_size = @as(u4, 1) << descriptor.content_size_flag;
468 const content_size = try in.takeVarInt(u64, .little, field_size);
469 break :c if (field_size == 2) content_size + 256 else content_size;
470 } else null;
471
472 return .{
473 .descriptor = descriptor,
474 .window_descriptor = window_descriptor,
475 .dictionary_id = dictionary_id,
476 .content_size = content_size,
477 };
478 }
479
480 /// Returns the window size required to decompress a frame, or `null` if it
481 /// cannot be determined (which indicates a malformed frame header).
482 pub fn windowSize(header: Header) ?u64 {
483 if (header.window_descriptor) |descriptor| {
484 const exponent = (descriptor & 0b11111000) >> 3;
485 const mantissa = descriptor & 0b00000111;
486 const window_log = 10 + exponent;
487 const window_base = @as(u64, 1) << @as(u6, @intCast(window_log));
488 const window_add = (window_base / 8) * mantissa;
489 return window_base + window_add;
490 } else return header.content_size;
491 }
492 };
493
494 pub const Block = struct {
495 pub const Header = packed struct(u24) {
496 last: bool,
497 type: Type,
498 size: u21,
499 };
500
501 pub const Type = enum(u2) {
502 raw,
503 rle,
504 compressed,
505 reserved,
506 };
507 };
508
509 pub const Decode = struct {
510 repeat_offsets: [3]u32,
511
512 offset: StateData(8),
513 match: StateData(9),
514 literal: StateData(9),
515
516 literal_fse_buffer: [zstd.table_size_max.literal]Table.Fse,
517 match_fse_buffer: [zstd.table_size_max.match]Table.Fse,
518 offset_fse_buffer: [zstd.table_size_max.offset]Table.Fse,
519
520 fse_tables_undefined: bool,
521
522 literal_stream_reader: ReverseBitReader,
523 literal_stream_index: usize,
524 literal_streams: LiteralsSection.Streams,
525 literal_header: LiteralsSection.Header,
526 huffman_tree: ?LiteralsSection.HuffmanTree,
527
528 literal_written_count: usize,
529
530 fn StateData(comptime max_accuracy_log: comptime_int) type {
531 return struct {
532 state: @This().State,
533 table: Table,
534 accuracy_log: u8,
535
536 const State = @Int(.unsigned, max_accuracy_log);
537 };
538 }
539
540 const init: Decode = .{
541 .repeat_offsets = .{
542 zstd.start_repeated_offset_1,
543 zstd.start_repeated_offset_2,
544 zstd.start_repeated_offset_3,
545 },
546
547 .offset = undefined,
548 .match = undefined,
549 .literal = undefined,
550
551 .literal_fse_buffer = undefined,
552 .match_fse_buffer = undefined,
553 .offset_fse_buffer = undefined,
554
555 .fse_tables_undefined = true,
556
557 .literal_written_count = 0,
558 .literal_header = undefined,
559 .literal_streams = undefined,
560 .literal_stream_reader = undefined,
561 .literal_stream_index = undefined,
562 .huffman_tree = null,
563 };
564
565 pub const PrepareError = error{
566 /// the (reversed) literal bitstream's first byte does not have any bits set
567 MissingStartBit,
568 /// `literals` is a treeless literals section and the decode state does not
569 /// have a Huffman tree from a previous block
570 TreelessLiteralsFirst,
571 /// on the first call if one of the sequence FSE tables is set to repeat mode
572 RepeatModeFirst,
573 /// an FSE table has an invalid accuracy
574 MalformedAccuracyLog,
575 /// failed decoding an FSE table
576 MalformedFseTable,
577 /// input stream ends before all FSE tables are read
578 EndOfStream,
579 ReadFailed,
580 InputBufferUndersize,
581 };
582
583 /// Prepare the decoder to decode a compressed block. Loads the
584 /// literals stream and Huffman tree from `literals` and reads the
585 /// FSE tables from `in`.
586 pub fn prepare(
587 self: *Decode,
588 in: *Reader,
589 remaining: *Limit,
590 literals: LiteralsSection,
591 sequences_header: SequencesSection.Header,
592 ) PrepareError!void {
593 self.literal_written_count = 0;
594 self.literal_header = literals.header;
595 self.literal_streams = literals.streams;
596
597 if (literals.huffman_tree) |tree| {
598 self.huffman_tree = tree;
599 } else if (literals.header.block_type == .treeless and self.huffman_tree == null) {
600 return error.TreelessLiteralsFirst;
601 }
602
603 switch (literals.header.block_type) {
604 .raw, .rle => {},
605 .compressed, .treeless => {
606 self.literal_stream_index = 0;
607 switch (literals.streams) {
608 .one => |slice| try self.initLiteralStream(slice),
609 .four => |streams| try self.initLiteralStream(streams[0]),
610 }
611 },
612 }
613
614 if (sequences_header.sequence_count > 0) {
615 try self.updateFseTable(in, remaining, .literal, sequences_header.literal_lengths);
616 try self.updateFseTable(in, remaining, .offset, sequences_header.offsets);
617 try self.updateFseTable(in, remaining, .match, sequences_header.match_lengths);
618 self.fse_tables_undefined = false;
619 }
620 }
621
622 /// Read initial FSE states for sequence decoding.
623 pub fn readInitialFseState(self: *Decode, bit_reader: *ReverseBitReader) error{EndOfStream}!void {
624 self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log);
625 self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log);
626 self.match.state = try bit_reader.readBitsNoEof(u9, self.match.accuracy_log);
627 }
628
629 fn updateRepeatOffset(self: *Decode, offset: u32) void {
630 self.repeat_offsets[2] = self.repeat_offsets[1];
631 self.repeat_offsets[1] = self.repeat_offsets[0];
632 self.repeat_offsets[0] = offset;
633 }
634
635 fn useRepeatOffset(self: *Decode, index: usize) u32 {
636 if (index == 1)
637 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[1])
638 else if (index == 2) {
639 std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[2]);
640 std.mem.swap(u32, &self.repeat_offsets[1], &self.repeat_offsets[2]);
641 }
642 return self.repeat_offsets[0];
643 }
644
645 const WhichFse = enum { offset, match, literal };
646
647 /// TODO: don't use `@field`
648 fn updateState(
649 self: *Decode,
650 comptime choice: WhichFse,
651 bit_reader: *ReverseBitReader,
652 ) error{ MalformedFseBits, EndOfStream }!void {
653 switch (@field(self, @tagName(choice)).table) {
654 .rle => {},
655 .fse => |table| {
656 const data = table[@field(self, @tagName(choice)).state];
657 const T = @TypeOf(@field(self, @tagName(choice))).State;
658 const bits_summand = try bit_reader.readBitsNoEof(T, data.bits);
659 const next_state = std.math.cast(
660 @TypeOf(@field(self, @tagName(choice))).State,
661 data.baseline + bits_summand,
662 ) orelse return error.MalformedFseBits;
663 @field(self, @tagName(choice)).state = next_state;
664 },
665 }
666 }
667
668 const FseTableError = error{
669 MalformedFseTable,
670 MalformedAccuracyLog,
671 RepeatModeFirst,
672 EndOfStream,
673 };
674
675 /// TODO: don't use `@field`
676 fn updateFseTable(
677 self: *Decode,
678 in: *Reader,
679 remaining: *Limit,
680 comptime choice: WhichFse,
681 mode: SequencesSection.Header.Mode,
682 ) !void {
683 const field_name = @tagName(choice);
684 switch (mode) {
685 .predefined => {
686 @field(self, field_name).accuracy_log =
687 @field(zstd.default_accuracy_log, field_name);
688
689 @field(self, field_name).table =
690 @field(Table, "predefined_" ++ field_name);
691 },
692 .rle => {
693 @field(self, field_name).accuracy_log = 0;
694 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
695 @field(self, field_name).table = .{ .rle = try in.takeByte() };
696 },
697 .fse => {
698 const max_table_size = 2048;
699 const peek_len: usize = remaining.minInt(max_table_size);
700 if (in.buffer.len < peek_len) return error.InputBufferUndersize;
701 const limited_buffer = try in.peek(peek_len);
702 var bit_reader: BitReader = .{ .bytes = limited_buffer };
703 const table_size = try Table.decode(
704 &bit_reader,
705 @field(zstd.table_symbol_count_max, field_name),
706 @field(zstd.table_accuracy_log_max, field_name),
707 &@field(self, field_name ++ "_fse_buffer"),
708 );
709 @field(self, field_name).table = .{
710 .fse = (&@field(self, field_name ++ "_fse_buffer"))[0..table_size],
711 };
712 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);
713 in.toss(bit_reader.index);
714 remaining.* = remaining.subtract(bit_reader.index).?;
715 },
716 .repeat => if (self.fse_tables_undefined) return error.RepeatModeFirst,
717 }
718 }
719
720 const Sequence = struct {
721 literal_length: u32,
722 match_length: u32,
723 offset: u32,
724 };
725
726 fn nextSequence(
727 self: *Decode,
728 bit_reader: *ReverseBitReader,
729 ) error{ InvalidBitStream, EndOfStream }!Sequence {
730 const raw_code = self.getCode(.offset);
731 const offset_code = std.math.cast(u5, raw_code) orelse {
732 return error.InvalidBitStream;
733 };
734 const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code);
735
736 const match_code = self.getCode(.match);
737 if (match_code >= zstd.match_length_code_table.len)
738 return error.InvalidBitStream;
739 const match = zstd.match_length_code_table[match_code];
740 const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]);
741
742 const literal_code = self.getCode(.literal);
743 if (literal_code >= zstd.literals_length_code_table.len)
744 return error.InvalidBitStream;
745 const literal = zstd.literals_length_code_table[literal_code];
746 const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]);
747
748 const offset = if (offset_value > 3) offset: {
749 const offset = offset_value - 3;
750 self.updateRepeatOffset(offset);
751 break :offset offset;
752 } else offset: {
753 if (literal_length == 0) {
754 if (offset_value == 3) {
755 const offset = self.repeat_offsets[0] - 1;
756 self.updateRepeatOffset(offset);
757 break :offset offset;
758 }
759 break :offset self.useRepeatOffset(offset_value);
760 }
761 break :offset self.useRepeatOffset(offset_value - 1);
762 };
763
764 if (offset == 0) return error.InvalidBitStream;
765
766 return .{
767 .literal_length = literal_length,
768 .match_length = match_length,
769 .offset = offset,
770 };
771 }
772
773 /// Decode one sequence from `bit_reader` into `dest`. Updates FSE states
774 /// if `last_sequence` is `false`. Assumes `prepare` called for the block
775 /// before attempting to decode sequences.
776 fn decodeSequence(
777 decode: *Decode,
778 dest: []u8,
779 write_pos: usize,
780 bit_reader: *ReverseBitReader,
781 ) !usize {
782 const sequence = try decode.nextSequence(bit_reader);
783 const literal_length: usize = sequence.literal_length;
784 const match_length: usize = sequence.match_length;
785 const sequence_length = literal_length + match_length;
786
787 if (sequence_length > dest[write_pos..].len)
788 return error.MalformedSequence;
789
790 const copy_start = std.math.sub(usize, write_pos + sequence.literal_length, sequence.offset) catch
791 return error.MalformedSequence;
792
793 if (decode.literal_written_count + literal_length > decode.literal_header.regenerated_size)
794 return error.MalformedLiteralsLength;
795 var sub_bw: Writer = .fixed(dest[write_pos..]);
796 try decodeLiterals(decode, &sub_bw, literal_length);
797 decode.literal_written_count += literal_length;
798 // This is not a @memmove; it intentionally repeats patterns
799 // caused by iterating one byte at a time.
800 for (
801 dest[write_pos + literal_length ..][0..match_length],
802 dest[copy_start..][0..match_length],
803 ) |*d, s| d.* = s;
804 return sequence_length;
805 }
806
807 fn nextLiteralMultiStream(self: *Decode) error{MissingStartBit}!void {
808 self.literal_stream_index += 1;
809 try self.initLiteralStream(self.literal_streams.four[self.literal_stream_index]);
810 }
811
812 fn initLiteralStream(self: *Decode, bytes: []const u8) error{MissingStartBit}!void {
813 self.literal_stream_reader = try ReverseBitReader.init(bytes);
814 }
815
816 fn isLiteralStreamEmpty(self: *Decode) bool {
817 switch (self.literal_streams) {
818 .one => return self.literal_stream_reader.isEmpty(),
819 .four => return self.literal_stream_index == 3 and self.literal_stream_reader.isEmpty(),
820 }
821 }
822
823 const LiteralBitsError = error{
824 MissingStartBit,
825 UnexpectedEndOfLiteralStream,
826 };
827 fn readLiteralsBits(
828 self: *Decode,
829 bit_count_to_read: u16,
830 ) LiteralBitsError!u16 {
831 return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: {
832 if (self.literal_streams == .four and self.literal_stream_index < 3) {
833 try self.nextLiteralMultiStream();
834 break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch
835 return error.UnexpectedEndOfLiteralStream;
836 } else {
837 return error.UnexpectedEndOfLiteralStream;
838 }
839 };
840 }
841
842 /// Decode `len` bytes of literals into `w`.
843 fn decodeLiterals(d: *Decode, w: *Writer, len: usize) !void {
844 switch (d.literal_header.block_type) {
845 .raw => {
846 try w.writeAll(d.literal_streams.one[d.literal_written_count..][0..len]);
847 },
848 .rle => {
849 try w.splatByteAll(d.literal_streams.one[0], len);
850 },
851 .compressed, .treeless => {
852 const buf = try w.writableSlice(len);
853 const huffman_tree = d.huffman_tree.?;
854 const max_bit_count = huffman_tree.max_bit_count;
855 const starting_bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
856 huffman_tree.nodes[huffman_tree.symbol_count_minus_one].weight,
857 max_bit_count,
858 );
859 var bits_read: u4 = 0;
860 var huffman_tree_index: usize = huffman_tree.symbol_count_minus_one;
861 var bit_count_to_read: u4 = starting_bit_count;
862 for (buf) |*out| {
863 var prefix: u16 = 0;
864 while (true) {
865 const new_bits = try d.readLiteralsBits(bit_count_to_read);
866 prefix <<= bit_count_to_read;
867 prefix |= new_bits;
868 bits_read += bit_count_to_read;
869 const result = try huffman_tree.query(huffman_tree_index, prefix);
870
871 switch (result) {
872 .symbol => |sym| {
873 out.* = sym;
874 bit_count_to_read = starting_bit_count;
875 bits_read = 0;
876 huffman_tree_index = huffman_tree.symbol_count_minus_one;
877 break;
878 },
879 .index => |index| {
880 huffman_tree_index = index;
881 const bit_count = LiteralsSection.HuffmanTree.weightToBitCount(
882 huffman_tree.nodes[index].weight,
883 max_bit_count,
884 );
885 bit_count_to_read = bit_count - bits_read;
886 },
887 }
888 }
889 }
890 },
891 }
892 }
893
894 /// TODO: don't use `@field`
895 fn getCode(self: *Decode, comptime choice: WhichFse) u32 {
896 return switch (@field(self, @tagName(choice)).table) {
897 .rle => |value| value,
898 .fse => |table| table[@field(self, @tagName(choice)).state].symbol,
899 };
900 }
901 };
902 };
903
904 pub const Skippable = struct {
905 pub const magic_min: Magic = @fromBackingInt(@intCast(0x184D2A50));
906 pub const magic_max: Magic = @fromBackingInt(@intCast(0x184D2A5F));
907
908 pub const Header = struct {
909 magic_number: u32,
910 frame_size: u32,
911 };
912 };
913
914 const InitError = error{
915 /// Frame uses a dictionary.
916 DictionaryIdFlagUnsupported,
917 /// Frame does not have a valid window size.
918 WindowSizeUnknown,
919 /// Window size exceeds `window_size_max` or max `usize` value.
920 WindowOversize,
921 /// Frame header indicates a content size exceeding max `usize` value.
922 ContentOversize,
923 };
924
925 /// Validates `frame_header` and returns the associated `Frame`.
926 pub fn init(
927 frame_header: Frame.Zstandard.Header,
928 window_size_max: usize,
929 verify_checksum: bool,
930 ) InitError!Frame {
931 if (frame_header.descriptor.dictionary_id_flag != 0)
932 return error.DictionaryIdFlagUnsupported;
933
934 const window_size_raw = frame_header.windowSize() orelse return error.WindowSizeUnknown;
935 const window_size = if (window_size_raw > window_size_max)
936 return error.WindowOversize
937 else
938 std.math.cast(usize, window_size_raw) orelse return error.WindowOversize;
939
940 const should_compute_checksum =
941 frame_header.descriptor.content_checksum_flag and verify_checksum;
942
943 const content_size = if (frame_header.content_size) |size|
944 std.math.cast(usize, size) orelse return error.ContentOversize
945 else
946 null;
947
948 return .{
949 .hasher_opt = if (should_compute_checksum) std.hash.XxHash64.init(0) else null,
950 .window_size = window_size,
951 .has_checksum = frame_header.descriptor.content_checksum_flag,
952 .block_size_max = @min(zstd.block_size_max, window_size),
953 .content_size = content_size,
954 };
955 }
956};
957
958pub const LiteralsSection = struct {
959 header: Header,
960 huffman_tree: ?HuffmanTree,
961 streams: Streams,
962
963 pub const Streams = union(enum) {
964 one: []const u8,
965 four: [4][]const u8,
966
967 fn decode(size_format: u2, stream_data: []const u8) !Streams {
968 if (size_format == 0) {
969 return .{ .one = stream_data };
970 }
971
972 if (stream_data.len < 6) return error.MalformedLiteralsSection;
973
974 const stream_1_length: usize = std.mem.readInt(u16, stream_data[0..2], .little);
975 const stream_2_length: usize = std.mem.readInt(u16, stream_data[2..4], .little);
976 const stream_3_length: usize = std.mem.readInt(u16, stream_data[4..6], .little);
977
978 const stream_1_start = 6;
979 const stream_2_start = stream_1_start + stream_1_length;
980 const stream_3_start = stream_2_start + stream_2_length;
981 const stream_4_start = stream_3_start + stream_3_length;
982
983 if (stream_data.len < stream_4_start) return error.MalformedLiteralsSection;
984
985 return .{ .four = .{
986 stream_data[stream_1_start .. stream_1_start + stream_1_length],
987 stream_data[stream_2_start .. stream_2_start + stream_2_length],
988 stream_data[stream_3_start .. stream_3_start + stream_3_length],
989 stream_data[stream_4_start..],
990 } };
991 }
992 };
993
994 pub const Header = struct {
995 block_type: BlockType,
996 size_format: u2,
997 regenerated_size: u20,
998 compressed_size: ?u18,
999
1000 /// Decode a literals section header.
1001 pub fn decode(in: *Reader, remaining: *Limit) !Header {
1002 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1003 const byte0 = try in.takeByte();
1004 const block_type: BlockType = @fromBackingInt(@intCast(byte0 & 0b11));
1005 const size_format: u2 = @intCast((byte0 & 0b1100) >> 2);
1006 var regenerated_size: u20 = undefined;
1007 var compressed_size: ?u18 = null;
1008 switch (block_type) {
1009 .raw, .rle => {
1010 switch (size_format) {
1011 0, 2 => {
1012 regenerated_size = byte0 >> 3;
1013 },
1014 1 => {
1015 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1016 regenerated_size = (byte0 >> 4) + (@as(u20, try in.takeByte()) << 4);
1017 },
1018 3 => {
1019 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
1020 regenerated_size = (byte0 >> 4) +
1021 (@as(u20, try in.takeByte()) << 4) +
1022 (@as(u20, try in.takeByte()) << 12);
1023 },
1024 }
1025 },
1026 .compressed, .treeless => {
1027 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
1028 const byte1 = try in.takeByte();
1029 const byte2 = try in.takeByte();
1030 switch (size_format) {
1031 0, 1 => {
1032 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);
1033 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);
1034 },
1035 2 => {
1036 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1037 const byte3 = try in.takeByte();
1038 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);
1039 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);
1040 },
1041 3 => {
1042 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
1043 const byte3 = try in.takeByte();
1044 const byte4 = try in.takeByte();
1045 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);
1046 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);
1047 },
1048 }
1049 },
1050 }
1051 return .{
1052 .block_type = block_type,
1053 .size_format = size_format,
1054 .regenerated_size = regenerated_size,
1055 .compressed_size = compressed_size,
1056 };
1057 }
1058 };
1059
1060 pub const BlockType = enum(u2) {
1061 raw,
1062 rle,
1063 compressed,
1064 treeless,
1065 };
1066
1067 pub const HuffmanTree = struct {
1068 max_bit_count: u4,
1069 symbol_count_minus_one: u8,
1070 nodes: [256]PrefixedSymbol,
1071
1072 pub const PrefixedSymbol = struct {
1073 symbol: u8,
1074 prefix: u16,
1075 weight: u4,
1076 };
1077
1078 pub const Result = union(enum) {
1079 symbol: u8,
1080 index: usize,
1081 };
1082
1083 pub fn query(self: HuffmanTree, index: usize, prefix: u16) error{HuffmanTreeIncomplete}!Result {
1084 var node = self.nodes[index];
1085 const weight = node.weight;
1086 var i: usize = index;
1087 while (node.weight == weight) {
1088 if (node.prefix == prefix) return .{ .symbol = node.symbol };
1089 if (i == 0) return error.HuffmanTreeIncomplete;
1090 i -= 1;
1091 node = self.nodes[i];
1092 }
1093 return .{ .index = i };
1094 }
1095
1096 pub fn weightToBitCount(weight: u4, max_bit_count: u4) u4 {
1097 return if (weight == 0) 0 else ((max_bit_count + 1) - weight);
1098 }
1099
1100 pub const DecodeError = Reader.Error || error{
1101 MalformedHuffmanTree,
1102 MalformedFseTable,
1103 MalformedAccuracyLog,
1104 EndOfStream,
1105 MissingStartBit,
1106 };
1107
1108 pub fn decode(in: *Reader, remaining: *Limit) HuffmanTree.DecodeError!HuffmanTree {
1109 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1110 const header = try in.takeByte();
1111 if (header < 128) {
1112 return decodeFse(in, remaining, header);
1113 } else {
1114 return decodeDirect(in, remaining, header - 127);
1115 }
1116 }
1117
1118 fn decodeDirect(
1119 in: *Reader,
1120 remaining: *Limit,
1121 encoded_symbol_count: usize,
1122 ) HuffmanTree.DecodeError!HuffmanTree {
1123 var weights: [256]u4 = undefined;
1124 const weights_byte_count = (encoded_symbol_count + 1) / 2;
1125 remaining.* = remaining.subtract(weights_byte_count) orelse return error.EndOfStream;
1126 for (0..weights_byte_count) |i| {
1127 const byte = try in.takeByte();
1128 weights[2 * i] = @as(u4, @intCast(byte >> 4));
1129 weights[2 * i + 1] = @as(u4, @intCast(byte & 0xF));
1130 }
1131 const symbol_count = encoded_symbol_count + 1;
1132 return build(&weights, symbol_count);
1133 }
1134
1135 fn decodeFse(
1136 in: *Reader,
1137 remaining: *Limit,
1138 compressed_size: usize,
1139 ) HuffmanTree.DecodeError!HuffmanTree {
1140 var weights: [256]u4 = undefined;
1141 remaining.* = remaining.subtract(compressed_size) orelse return error.EndOfStream;
1142 const compressed_buffer = try in.take(compressed_size);
1143 var bit_reader: BitReader = .{ .bytes = compressed_buffer };
1144 var entries: [1 << 6]Table.Fse = undefined;
1145 const table_size = try Table.decode(&bit_reader, 256, 6, &entries);
1146 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
1147 const remaining_buffer = bit_reader.bytes[bit_reader.index..];
1148 const symbol_count = try assignWeights(remaining_buffer, accuracy_log, &entries, &weights);
1149 return build(&weights, symbol_count);
1150 }
1151
1152 fn assignWeights(
1153 huff_bits_buffer: []const u8,
1154 accuracy_log: u16,
1155 entries: *[1 << 6]Table.Fse,
1156 weights: *[256]u4,
1157 ) !usize {
1158 var huff_bits = try ReverseBitReader.init(huff_bits_buffer);
1159
1160 var i: usize = 0;
1161 var even_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log);
1162 var odd_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log);
1163
1164 while (i < 254) {
1165 const even_data = entries[even_state];
1166 var read_bits: u16 = 0;
1167 const even_bits = huff_bits.readBits(u32, even_data.bits, &read_bits) catch unreachable;
1168 weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree;
1169 i += 1;
1170 if (read_bits < even_data.bits) {
1171 weights[i] = std.math.cast(u4, entries[odd_state].symbol) orelse return error.MalformedHuffmanTree;
1172 i += 1;
1173 break;
1174 }
1175 even_state = even_data.baseline + even_bits;
1176
1177 read_bits = 0;
1178 const odd_data = entries[odd_state];
1179 const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable;
1180 weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree;
1181 i += 1;
1182 if (read_bits < odd_data.bits) {
1183 if (i == 255) return error.MalformedHuffmanTree;
1184 weights[i] = std.math.cast(u4, entries[even_state].symbol) orelse return error.MalformedHuffmanTree;
1185 i += 1;
1186 break;
1187 }
1188 odd_state = odd_data.baseline + odd_bits;
1189 } else return error.MalformedHuffmanTree;
1190
1191 if (!huff_bits.isEmpty()) {
1192 return error.MalformedHuffmanTree;
1193 }
1194
1195 return i + 1; // stream contains all but the last symbol
1196 }
1197
1198 fn assignSymbols(weight_sorted_prefixed_symbols: []PrefixedSymbol, weights: [256]u4) usize {
1199 for (0..weight_sorted_prefixed_symbols.len) |i| {
1200 weight_sorted_prefixed_symbols[i] = .{
1201 .symbol = @as(u8, @intCast(i)),
1202 .weight = undefined,
1203 .prefix = undefined,
1204 };
1205 }
1206
1207 std.mem.sort(
1208 PrefixedSymbol,
1209 weight_sorted_prefixed_symbols,
1210 weights,
1211 lessThanByWeight,
1212 );
1213
1214 var prefix: u16 = 0;
1215 var prefixed_symbol_count: usize = 0;
1216 var sorted_index: usize = 0;
1217 const symbol_count = weight_sorted_prefixed_symbols.len;
1218 while (sorted_index < symbol_count) {
1219 var symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
1220 const weight = weights[symbol];
1221 if (weight == 0) {
1222 sorted_index += 1;
1223 continue;
1224 }
1225
1226 while (sorted_index < symbol_count) : ({
1227 sorted_index += 1;
1228 prefixed_symbol_count += 1;
1229 prefix += 1;
1230 }) {
1231 symbol = weight_sorted_prefixed_symbols[sorted_index].symbol;
1232 if (weights[symbol] != weight) {
1233 prefix = ((prefix - 1) >> (weights[symbol] - weight)) + 1;
1234 break;
1235 }
1236 weight_sorted_prefixed_symbols[prefixed_symbol_count].symbol = symbol;
1237 weight_sorted_prefixed_symbols[prefixed_symbol_count].prefix = prefix;
1238 weight_sorted_prefixed_symbols[prefixed_symbol_count].weight = weight;
1239 }
1240 }
1241 return prefixed_symbol_count;
1242 }
1243
1244 fn build(weights: *[256]u4, symbol_count: usize) error{MalformedHuffmanTree}!HuffmanTree {
1245 var weight_power_sum_big: u32 = 0;
1246 for (weights[0 .. symbol_count - 1]) |value| {
1247 weight_power_sum_big += (@as(u16, 1) << value) >> 1;
1248 }
1249 if (weight_power_sum_big >= 1 << 11) return error.MalformedHuffmanTree;
1250 const weight_power_sum = @as(u16, @intCast(weight_power_sum_big));
1251
1252 // advance to next power of two (even if weight_power_sum is a power of 2)
1253 // TODO: is it valid to have weight_power_sum == 0?
1254 const max_number_of_bits = if (weight_power_sum == 0) 1 else std.math.log2_int(u16, weight_power_sum) + 1;
1255 const next_power_of_two = @as(u16, 1) << max_number_of_bits;
1256 weights[symbol_count - 1] = std.math.log2_int(u16, next_power_of_two - weight_power_sum) + 1;
1257
1258 var weight_sorted_prefixed_symbols: [256]PrefixedSymbol = undefined;
1259 const prefixed_symbol_count = assignSymbols(weight_sorted_prefixed_symbols[0..symbol_count], weights.*);
1260 const tree: HuffmanTree = .{
1261 .max_bit_count = max_number_of_bits,
1262 .symbol_count_minus_one = @as(u8, @intCast(prefixed_symbol_count - 1)),
1263 .nodes = weight_sorted_prefixed_symbols,
1264 };
1265 return tree;
1266 }
1267
1268 fn lessThanByWeight(
1269 weights: [256]u4,
1270 lhs: PrefixedSymbol,
1271 rhs: PrefixedSymbol,
1272 ) bool {
1273 // NOTE: this function relies on the use of a stable sorting algorithm,
1274 // otherwise a special case of if (weights[lhs] == weights[rhs]) return lhs < rhs;
1275 // should be added
1276 return weights[lhs.symbol] < weights[rhs.symbol];
1277 }
1278 };
1279
1280 pub const StreamCount = enum { one, four };
1281 pub fn streamCount(size_format: u2, block_type: BlockType) StreamCount {
1282 return switch (block_type) {
1283 .raw, .rle => .one,
1284 .compressed, .treeless => if (size_format == 0) .one else .four,
1285 };
1286 }
1287
1288 pub const DecodeError = error{
1289 /// Invalid header.
1290 MalformedLiteralsHeader,
1291 /// Decoding errors.
1292 MalformedLiteralsSection,
1293 /// Compressed literals have invalid accuracy.
1294 MalformedAccuracyLog,
1295 /// Compressed literals have invalid FSE table.
1296 MalformedFseTable,
1297 /// Failed decoding a Huffamn tree.
1298 MalformedHuffmanTree,
1299 /// Not enough bytes to complete the section.
1300 EndOfStream,
1301 ReadFailed,
1302 MissingStartBit,
1303 };
1304
1305 pub fn decode(in: *Reader, remaining: *Limit, buffer: []u8) DecodeError!LiteralsSection {
1306 const header = try Header.decode(in, remaining);
1307 switch (header.block_type) {
1308 .raw => {
1309 if (buffer.len < header.regenerated_size) return error.MalformedLiteralsSection;
1310 remaining.* = remaining.subtract(header.regenerated_size) orelse return error.EndOfStream;
1311 try in.readSliceAll(buffer[0..header.regenerated_size]);
1312 return .{
1313 .header = header,
1314 .huffman_tree = null,
1315 .streams = .{ .one = buffer },
1316 };
1317 },
1318 .rle => {
1319 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1320 buffer[0] = try in.takeByte();
1321 return .{
1322 .header = header,
1323 .huffman_tree = null,
1324 .streams = .{ .one = buffer[0..1] },
1325 };
1326 },
1327 .compressed, .treeless => {
1328 const before_remaining = remaining.*;
1329 const huffman_tree = if (header.block_type == .compressed)
1330 try HuffmanTree.decode(in, remaining)
1331 else
1332 null;
1333 const huffman_tree_size = @backingInt(before_remaining) - @backingInt(remaining.*);
1334 const total_streams_size = std.math.sub(usize, header.compressed_size.?, huffman_tree_size) catch
1335 return error.MalformedLiteralsSection;
1336 if (total_streams_size > buffer.len) return error.MalformedLiteralsSection;
1337 remaining.* = remaining.subtract(total_streams_size) orelse return error.EndOfStream;
1338 try in.readSliceAll(buffer[0..total_streams_size]);
1339 const stream_data = buffer[0..total_streams_size];
1340 const streams = try Streams.decode(header.size_format, stream_data);
1341 return .{
1342 .header = header,
1343 .huffman_tree = huffman_tree,
1344 .streams = streams,
1345 };
1346 },
1347 }
1348 }
1349};
1350
1351pub const SequencesSection = struct {
1352 header: Header,
1353 literals_length_table: Table,
1354 offset_table: Table,
1355 match_length_table: Table,
1356
1357 pub const Header = struct {
1358 sequence_count: u24,
1359 match_lengths: Mode,
1360 offsets: Mode,
1361 literal_lengths: Mode,
1362
1363 pub const Mode = enum(u2) {
1364 predefined,
1365 rle,
1366 fse,
1367 repeat,
1368 };
1369
1370 pub const DecodeError = error{
1371 ReservedBitSet,
1372 EndOfStream,
1373 ReadFailed,
1374 };
1375
1376 pub fn decode(in: *Reader, remaining: *Limit) DecodeError!Header {
1377 var sequence_count: u24 = undefined;
1378
1379 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1380 const byte0 = try in.takeByte();
1381 if (byte0 == 0) {
1382 return .{
1383 .sequence_count = 0,
1384 .offsets = undefined,
1385 .match_lengths = undefined,
1386 .literal_lengths = undefined,
1387 };
1388 } else if (byte0 < 128) {
1389 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1390 sequence_count = byte0;
1391 } else if (byte0 < 255) {
1392 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
1393 sequence_count = (@as(u24, (byte0 - 128)) << 8) + try in.takeByte();
1394 } else {
1395 remaining.* = remaining.subtract(3) orelse return error.EndOfStream;
1396 sequence_count = (try in.takeByte()) + (@as(u24, try in.takeByte()) << 8) + 0x7F00;
1397 }
1398
1399 const compression_modes = try in.takeByte();
1400
1401 const matches_mode: Header.Mode = @fromBackingInt(@intCast((compression_modes & 0b00001100) >> 2));
1402 const offsets_mode: Header.Mode = @fromBackingInt(@intCast((compression_modes & 0b00110000) >> 4));
1403 const literal_mode: Header.Mode = @fromBackingInt(@intCast((compression_modes & 0b11000000) >> 6));
1404 if (compression_modes & 0b11 != 0) return error.ReservedBitSet;
1405
1406 return .{
1407 .sequence_count = sequence_count,
1408 .offsets = offsets_mode,
1409 .match_lengths = matches_mode,
1410 .literal_lengths = literal_mode,
1411 };
1412 }
1413 };
1414};
1415
1416pub const Table = union(enum) {
1417 fse: []const Fse,
1418 rle: u8,
1419
1420 pub const Fse = struct {
1421 symbol: u8,
1422 baseline: u16,
1423 bits: u8,
1424 };
1425
1426 pub fn decode(
1427 bit_reader: *BitReader,
1428 expected_symbol_count: usize,
1429 max_accuracy_log: u4,
1430 entries: []Table.Fse,
1431 ) !usize {
1432 const accuracy_log_biased = try bit_reader.readBitsNoEof(u4, 4);
1433 if (accuracy_log_biased > max_accuracy_log -| 5) return error.MalformedAccuracyLog;
1434 const accuracy_log = accuracy_log_biased + 5;
1435
1436 var values: [256]u16 = undefined;
1437 var value_count: usize = 0;
1438
1439 const total_probability = @as(u16, 1) << accuracy_log;
1440 var accumulated_probability: u16 = 0;
1441
1442 while (accumulated_probability < total_probability) {
1443 // WARNING: The RFC is poorly worded, and would suggest std.math.log2_int_ceil is correct here,
1444 // but power of two (remaining probabilities + 1) need max bits set to 1 more.
1445 const max_bits = std.math.log2_int(u16, total_probability - accumulated_probability + 1) + 1;
1446 const small = try bit_reader.readBitsNoEof(u16, max_bits - 1);
1447
1448 const cutoff = (@as(u16, 1) << max_bits) - 1 - (total_probability - accumulated_probability + 1);
1449
1450 const value = if (small < cutoff)
1451 small
1452 else value: {
1453 const value_read = small + (try bit_reader.readBitsNoEof(u16, 1) << (max_bits - 1));
1454 break :value if (value_read < @as(u16, 1) << (max_bits - 1))
1455 value_read
1456 else
1457 value_read - cutoff;
1458 };
1459
1460 accumulated_probability += if (value != 0) value - 1 else 1;
1461
1462 values[value_count] = value;
1463 value_count += 1;
1464
1465 if (value == 1) {
1466 while (true) {
1467 const repeat_flag = try bit_reader.readBitsNoEof(u2, 2);
1468 if (repeat_flag + value_count > 256) return error.MalformedFseTable;
1469 for (0..repeat_flag) |_| {
1470 values[value_count] = 1;
1471 value_count += 1;
1472 }
1473 if (repeat_flag < 3) break;
1474 }
1475 }
1476 if (value_count == 256) break;
1477 }
1478 bit_reader.alignToByte();
1479
1480 if (value_count < 2) return error.MalformedFseTable;
1481 if (accumulated_probability != total_probability) return error.MalformedFseTable;
1482 if (value_count > expected_symbol_count) return error.MalformedFseTable;
1483
1484 const table_size = total_probability;
1485
1486 try build(values[0..value_count], entries[0..table_size]);
1487 return table_size;
1488 }
1489
1490 pub fn build(values: []const u16, entries: []Table.Fse) !void {
1491 const total_probability = @as(u16, @intCast(entries.len));
1492 const accuracy_log = std.math.log2_int(u16, total_probability);
1493 assert(total_probability <= 1 << 9);
1494
1495 var less_than_one_count: usize = 0;
1496 for (values, 0..) |value, i| {
1497 if (value == 0) {
1498 entries[entries.len - 1 - less_than_one_count] = Table.Fse{
1499 .symbol = @as(u8, @intCast(i)),
1500 .baseline = 0,
1501 .bits = accuracy_log,
1502 };
1503 less_than_one_count += 1;
1504 }
1505 }
1506
1507 var position: usize = 0;
1508 var temp_states: [1 << 9]u16 = undefined;
1509 for (values, 0..) |value, symbol| {
1510 if (value == 0 or value == 1) continue;
1511 const probability = value - 1;
1512
1513 const state_share_dividend = std.math.ceilPowerOfTwo(u16, probability) catch
1514 return error.MalformedFseTable;
1515 const share_size = @divExact(total_probability, state_share_dividend);
1516 const double_state_count = state_share_dividend - probability;
1517 const single_state_count = probability - double_state_count;
1518 const share_size_log = std.math.log2_int(u16, share_size);
1519
1520 for (0..probability) |i| {
1521 temp_states[i] = @as(u16, @intCast(position));
1522 position += (entries.len >> 1) + (entries.len >> 3) + 3;
1523 position &= entries.len - 1;
1524 while (position >= entries.len - less_than_one_count) {
1525 position += (entries.len >> 1) + (entries.len >> 3) + 3;
1526 position &= entries.len - 1;
1527 }
1528 }
1529 std.mem.sort(u16, temp_states[0..probability], {}, std.sort.asc(u16));
1530 for (0..probability) |i| {
1531 entries[temp_states[i]] = if (i < double_state_count) Table.Fse{
1532 .symbol = @as(u8, @intCast(symbol)),
1533 .bits = share_size_log + 1,
1534 .baseline = single_state_count * share_size + @as(u16, @intCast(i)) * 2 * share_size,
1535 } else Table.Fse{
1536 .symbol = @as(u8, @intCast(symbol)),
1537 .bits = share_size_log,
1538 .baseline = (@as(u16, @intCast(i)) - double_state_count) * share_size,
1539 };
1540 }
1541 }
1542 }
1543
1544 test build {
1545 const literals_length_default_values = [36]u16{
1546 5, 4, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2,
1547 3, 3, 3, 3, 3, 3, 3, 3, 3, 4, 3, 2, 2, 2, 2, 2,
1548 0, 0, 0, 0,
1549 };
1550
1551 const match_lengths_default_values = [53]u16{
1552 2, 5, 4, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2,
1553 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
1554 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 0, 0,
1555 0, 0, 0, 0, 0,
1556 };
1557
1558 const offset_codes_default_values = [29]u16{
1559 2, 2, 2, 2, 2, 2, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2,
1560 2, 2, 2, 2, 2, 2, 2, 2, 0, 0, 0, 0, 0,
1561 };
1562
1563 var entries: [64]Table.Fse = undefined;
1564 try build(&literals_length_default_values, &entries);
1565 try std.testing.expectEqualSlices(Table.Fse, Table.predefined_literal.fse, &entries);
1566
1567 try build(&match_lengths_default_values, &entries);
1568 try std.testing.expectEqualSlices(Table.Fse, Table.predefined_match.fse, &entries);
1569
1570 try build(&offset_codes_default_values, entries[0..32]);
1571 try std.testing.expectEqualSlices(Table.Fse, Table.predefined_offset.fse, entries[0..32]);
1572 }
1573
1574 pub const predefined_literal: Table = .{
1575 .fse = &[64]Table.Fse{
1576 .{ .symbol = 0, .bits = 4, .baseline = 0 },
1577 .{ .symbol = 0, .bits = 4, .baseline = 16 },
1578 .{ .symbol = 1, .bits = 5, .baseline = 32 },
1579 .{ .symbol = 3, .bits = 5, .baseline = 0 },
1580 .{ .symbol = 4, .bits = 5, .baseline = 0 },
1581 .{ .symbol = 6, .bits = 5, .baseline = 0 },
1582 .{ .symbol = 7, .bits = 5, .baseline = 0 },
1583 .{ .symbol = 9, .bits = 5, .baseline = 0 },
1584 .{ .symbol = 10, .bits = 5, .baseline = 0 },
1585 .{ .symbol = 12, .bits = 5, .baseline = 0 },
1586 .{ .symbol = 14, .bits = 6, .baseline = 0 },
1587 .{ .symbol = 16, .bits = 5, .baseline = 0 },
1588 .{ .symbol = 18, .bits = 5, .baseline = 0 },
1589 .{ .symbol = 19, .bits = 5, .baseline = 0 },
1590 .{ .symbol = 21, .bits = 5, .baseline = 0 },
1591 .{ .symbol = 22, .bits = 5, .baseline = 0 },
1592 .{ .symbol = 24, .bits = 5, .baseline = 0 },
1593 .{ .symbol = 25, .bits = 5, .baseline = 32 },
1594 .{ .symbol = 26, .bits = 5, .baseline = 0 },
1595 .{ .symbol = 27, .bits = 6, .baseline = 0 },
1596 .{ .symbol = 29, .bits = 6, .baseline = 0 },
1597 .{ .symbol = 31, .bits = 6, .baseline = 0 },
1598 .{ .symbol = 0, .bits = 4, .baseline = 32 },
1599 .{ .symbol = 1, .bits = 4, .baseline = 0 },
1600 .{ .symbol = 2, .bits = 5, .baseline = 0 },
1601 .{ .symbol = 4, .bits = 5, .baseline = 32 },
1602 .{ .symbol = 5, .bits = 5, .baseline = 0 },
1603 .{ .symbol = 7, .bits = 5, .baseline = 32 },
1604 .{ .symbol = 8, .bits = 5, .baseline = 0 },
1605 .{ .symbol = 10, .bits = 5, .baseline = 32 },
1606 .{ .symbol = 11, .bits = 5, .baseline = 0 },
1607 .{ .symbol = 13, .bits = 6, .baseline = 0 },
1608 .{ .symbol = 16, .bits = 5, .baseline = 32 },
1609 .{ .symbol = 17, .bits = 5, .baseline = 0 },
1610 .{ .symbol = 19, .bits = 5, .baseline = 32 },
1611 .{ .symbol = 20, .bits = 5, .baseline = 0 },
1612 .{ .symbol = 22, .bits = 5, .baseline = 32 },
1613 .{ .symbol = 23, .bits = 5, .baseline = 0 },
1614 .{ .symbol = 25, .bits = 4, .baseline = 0 },
1615 .{ .symbol = 25, .bits = 4, .baseline = 16 },
1616 .{ .symbol = 26, .bits = 5, .baseline = 32 },
1617 .{ .symbol = 28, .bits = 6, .baseline = 0 },
1618 .{ .symbol = 30, .bits = 6, .baseline = 0 },
1619 .{ .symbol = 0, .bits = 4, .baseline = 48 },
1620 .{ .symbol = 1, .bits = 4, .baseline = 16 },
1621 .{ .symbol = 2, .bits = 5, .baseline = 32 },
1622 .{ .symbol = 3, .bits = 5, .baseline = 32 },
1623 .{ .symbol = 5, .bits = 5, .baseline = 32 },
1624 .{ .symbol = 6, .bits = 5, .baseline = 32 },
1625 .{ .symbol = 8, .bits = 5, .baseline = 32 },
1626 .{ .symbol = 9, .bits = 5, .baseline = 32 },
1627 .{ .symbol = 11, .bits = 5, .baseline = 32 },
1628 .{ .symbol = 12, .bits = 5, .baseline = 32 },
1629 .{ .symbol = 15, .bits = 6, .baseline = 0 },
1630 .{ .symbol = 17, .bits = 5, .baseline = 32 },
1631 .{ .symbol = 18, .bits = 5, .baseline = 32 },
1632 .{ .symbol = 20, .bits = 5, .baseline = 32 },
1633 .{ .symbol = 21, .bits = 5, .baseline = 32 },
1634 .{ .symbol = 23, .bits = 5, .baseline = 32 },
1635 .{ .symbol = 24, .bits = 5, .baseline = 32 },
1636 .{ .symbol = 35, .bits = 6, .baseline = 0 },
1637 .{ .symbol = 34, .bits = 6, .baseline = 0 },
1638 .{ .symbol = 33, .bits = 6, .baseline = 0 },
1639 .{ .symbol = 32, .bits = 6, .baseline = 0 },
1640 },
1641 };
1642
1643 pub const predefined_match: Table = .{
1644 .fse = &[64]Table.Fse{
1645 .{ .symbol = 0, .bits = 6, .baseline = 0 },
1646 .{ .symbol = 1, .bits = 4, .baseline = 0 },
1647 .{ .symbol = 2, .bits = 5, .baseline = 32 },
1648 .{ .symbol = 3, .bits = 5, .baseline = 0 },
1649 .{ .symbol = 5, .bits = 5, .baseline = 0 },
1650 .{ .symbol = 6, .bits = 5, .baseline = 0 },
1651 .{ .symbol = 8, .bits = 5, .baseline = 0 },
1652 .{ .symbol = 10, .bits = 6, .baseline = 0 },
1653 .{ .symbol = 13, .bits = 6, .baseline = 0 },
1654 .{ .symbol = 16, .bits = 6, .baseline = 0 },
1655 .{ .symbol = 19, .bits = 6, .baseline = 0 },
1656 .{ .symbol = 22, .bits = 6, .baseline = 0 },
1657 .{ .symbol = 25, .bits = 6, .baseline = 0 },
1658 .{ .symbol = 28, .bits = 6, .baseline = 0 },
1659 .{ .symbol = 31, .bits = 6, .baseline = 0 },
1660 .{ .symbol = 33, .bits = 6, .baseline = 0 },
1661 .{ .symbol = 35, .bits = 6, .baseline = 0 },
1662 .{ .symbol = 37, .bits = 6, .baseline = 0 },
1663 .{ .symbol = 39, .bits = 6, .baseline = 0 },
1664 .{ .symbol = 41, .bits = 6, .baseline = 0 },
1665 .{ .symbol = 43, .bits = 6, .baseline = 0 },
1666 .{ .symbol = 45, .bits = 6, .baseline = 0 },
1667 .{ .symbol = 1, .bits = 4, .baseline = 16 },
1668 .{ .symbol = 2, .bits = 4, .baseline = 0 },
1669 .{ .symbol = 3, .bits = 5, .baseline = 32 },
1670 .{ .symbol = 4, .bits = 5, .baseline = 0 },
1671 .{ .symbol = 6, .bits = 5, .baseline = 32 },
1672 .{ .symbol = 7, .bits = 5, .baseline = 0 },
1673 .{ .symbol = 9, .bits = 6, .baseline = 0 },
1674 .{ .symbol = 12, .bits = 6, .baseline = 0 },
1675 .{ .symbol = 15, .bits = 6, .baseline = 0 },
1676 .{ .symbol = 18, .bits = 6, .baseline = 0 },
1677 .{ .symbol = 21, .bits = 6, .baseline = 0 },
1678 .{ .symbol = 24, .bits = 6, .baseline = 0 },
1679 .{ .symbol = 27, .bits = 6, .baseline = 0 },
1680 .{ .symbol = 30, .bits = 6, .baseline = 0 },
1681 .{ .symbol = 32, .bits = 6, .baseline = 0 },
1682 .{ .symbol = 34, .bits = 6, .baseline = 0 },
1683 .{ .symbol = 36, .bits = 6, .baseline = 0 },
1684 .{ .symbol = 38, .bits = 6, .baseline = 0 },
1685 .{ .symbol = 40, .bits = 6, .baseline = 0 },
1686 .{ .symbol = 42, .bits = 6, .baseline = 0 },
1687 .{ .symbol = 44, .bits = 6, .baseline = 0 },
1688 .{ .symbol = 1, .bits = 4, .baseline = 32 },
1689 .{ .symbol = 1, .bits = 4, .baseline = 48 },
1690 .{ .symbol = 2, .bits = 4, .baseline = 16 },
1691 .{ .symbol = 4, .bits = 5, .baseline = 32 },
1692 .{ .symbol = 5, .bits = 5, .baseline = 32 },
1693 .{ .symbol = 7, .bits = 5, .baseline = 32 },
1694 .{ .symbol = 8, .bits = 5, .baseline = 32 },
1695 .{ .symbol = 11, .bits = 6, .baseline = 0 },
1696 .{ .symbol = 14, .bits = 6, .baseline = 0 },
1697 .{ .symbol = 17, .bits = 6, .baseline = 0 },
1698 .{ .symbol = 20, .bits = 6, .baseline = 0 },
1699 .{ .symbol = 23, .bits = 6, .baseline = 0 },
1700 .{ .symbol = 26, .bits = 6, .baseline = 0 },
1701 .{ .symbol = 29, .bits = 6, .baseline = 0 },
1702 .{ .symbol = 52, .bits = 6, .baseline = 0 },
1703 .{ .symbol = 51, .bits = 6, .baseline = 0 },
1704 .{ .symbol = 50, .bits = 6, .baseline = 0 },
1705 .{ .symbol = 49, .bits = 6, .baseline = 0 },
1706 .{ .symbol = 48, .bits = 6, .baseline = 0 },
1707 .{ .symbol = 47, .bits = 6, .baseline = 0 },
1708 .{ .symbol = 46, .bits = 6, .baseline = 0 },
1709 },
1710 };
1711
1712 pub const predefined_offset: Table = .{
1713 .fse = &[32]Table.Fse{
1714 .{ .symbol = 0, .bits = 5, .baseline = 0 },
1715 .{ .symbol = 6, .bits = 4, .baseline = 0 },
1716 .{ .symbol = 9, .bits = 5, .baseline = 0 },
1717 .{ .symbol = 15, .bits = 5, .baseline = 0 },
1718 .{ .symbol = 21, .bits = 5, .baseline = 0 },
1719 .{ .symbol = 3, .bits = 5, .baseline = 0 },
1720 .{ .symbol = 7, .bits = 4, .baseline = 0 },
1721 .{ .symbol = 12, .bits = 5, .baseline = 0 },
1722 .{ .symbol = 18, .bits = 5, .baseline = 0 },
1723 .{ .symbol = 23, .bits = 5, .baseline = 0 },
1724 .{ .symbol = 5, .bits = 5, .baseline = 0 },
1725 .{ .symbol = 8, .bits = 4, .baseline = 0 },
1726 .{ .symbol = 14, .bits = 5, .baseline = 0 },
1727 .{ .symbol = 20, .bits = 5, .baseline = 0 },
1728 .{ .symbol = 2, .bits = 5, .baseline = 0 },
1729 .{ .symbol = 7, .bits = 4, .baseline = 16 },
1730 .{ .symbol = 11, .bits = 5, .baseline = 0 },
1731 .{ .symbol = 17, .bits = 5, .baseline = 0 },
1732 .{ .symbol = 22, .bits = 5, .baseline = 0 },
1733 .{ .symbol = 4, .bits = 5, .baseline = 0 },
1734 .{ .symbol = 8, .bits = 4, .baseline = 16 },
1735 .{ .symbol = 13, .bits = 5, .baseline = 0 },
1736 .{ .symbol = 19, .bits = 5, .baseline = 0 },
1737 .{ .symbol = 1, .bits = 5, .baseline = 0 },
1738 .{ .symbol = 6, .bits = 4, .baseline = 16 },
1739 .{ .symbol = 10, .bits = 5, .baseline = 0 },
1740 .{ .symbol = 16, .bits = 5, .baseline = 0 },
1741 .{ .symbol = 28, .bits = 5, .baseline = 0 },
1742 .{ .symbol = 27, .bits = 5, .baseline = 0 },
1743 .{ .symbol = 26, .bits = 5, .baseline = 0 },
1744 .{ .symbol = 25, .bits = 5, .baseline = 0 },
1745 .{ .symbol = 24, .bits = 5, .baseline = 0 },
1746 },
1747 };
1748};
1749
1750const low_bit_mask = [9]u8{
1751 0b00000000,
1752 0b00000001,
1753 0b00000011,
1754 0b00000111,
1755 0b00001111,
1756 0b00011111,
1757 0b00111111,
1758 0b01111111,
1759 0b11111111,
1760};
1761
1762fn Bits(comptime T: type) type {
1763 return struct { T, u16 };
1764}
1765
1766/// For reading the reversed bit streams used to encode FSE compressed data.
1767const ReverseBitReader = struct {
1768 bytes: []const u8,
1769 remaining: usize,
1770 bits: u8,
1771 count: u4,
1772
1773 fn init(bytes: []const u8) error{MissingStartBit}!ReverseBitReader {
1774 var result: ReverseBitReader = .{
1775 .bytes = bytes,
1776 .remaining = bytes.len,
1777 .bits = 0,
1778 .count = 0,
1779 };
1780 if (bytes.len == 0) return result;
1781 for (0..8) |_| if (0 != (result.readBitsNoEof(u1, 1) catch unreachable)) return result;
1782 return error.MissingStartBit;
1783 }
1784
1785 fn initBits(comptime T: type, out: anytype, num: u16) Bits(T) {
1786 const UT = @Int(.unsigned, @bitSizeOf(T));
1787 return .{
1788 @bitCast(@as(UT, @intCast(out))),
1789 num,
1790 };
1791 }
1792
1793 fn readBitsNoEof(self: *ReverseBitReader, comptime T: type, num: u16) error{EndOfStream}!T {
1794 const b, const c = try self.readBitsTuple(T, num);
1795 if (c < num) return error.EndOfStream;
1796 return b;
1797 }
1798
1799 fn readBits(self: *ReverseBitReader, comptime T: type, num: u16, out_bits: *u16) !T {
1800 const b, const c = try self.readBitsTuple(T, num);
1801 out_bits.* = c;
1802 return b;
1803 }
1804
1805 fn readBitsTuple(self: *ReverseBitReader, comptime T: type, num: u16) !Bits(T) {
1806 const UT = @Int(.unsigned, @bitSizeOf(T));
1807 const U = if (@bitSizeOf(T) < 8) u8 else UT;
1808
1809 if (num <= self.count) return initBits(T, self.removeBits(@intCast(num)), num);
1810
1811 var out_count: u16 = self.count;
1812 var out: U = self.removeBits(self.count);
1813
1814 const full_bytes_left = (num - out_count) / 8;
1815
1816 for (0..full_bytes_left) |_| {
1817 const byte = takeByte(self) catch |err| switch (err) {
1818 error.EndOfStream => return initBits(T, out, out_count),
1819 };
1820 if (U == u8) out = 0 else out <<= 8;
1821 out |= byte;
1822 out_count += 8;
1823 }
1824
1825 const bits_left = num - out_count;
1826 const keep = 8 - bits_left;
1827
1828 if (bits_left == 0) return initBits(T, out, out_count);
1829
1830 const final_byte = takeByte(self) catch |err| switch (err) {
1831 error.EndOfStream => return initBits(T, out, out_count),
1832 };
1833
1834 out <<= @intCast(bits_left);
1835 out |= final_byte >> @intCast(keep);
1836 self.bits = final_byte & low_bit_mask[keep];
1837
1838 self.count = @intCast(keep);
1839 return initBits(T, out, num);
1840 }
1841
1842 fn takeByte(rbr: *ReverseBitReader) error{EndOfStream}!u8 {
1843 if (rbr.remaining == 0) return error.EndOfStream;
1844 rbr.remaining -= 1;
1845 return rbr.bytes[rbr.remaining];
1846 }
1847
1848 fn isEmpty(self: *const ReverseBitReader) bool {
1849 return self.remaining == 0 and self.count == 0;
1850 }
1851
1852 fn removeBits(self: *ReverseBitReader, num: u4) u8 {
1853 if (num == 8) {
1854 self.count = 0;
1855 return self.bits;
1856 }
1857
1858 const keep = self.count - num;
1859 const bits = self.bits >> @intCast(keep);
1860 self.bits &= low_bit_mask[keep];
1861
1862 self.count = keep;
1863 return bits;
1864 }
1865};
1866
1867const BitReader = struct {
1868 bytes: []const u8,
1869 index: usize = 0,
1870 bits: u8 = 0,
1871 count: u4 = 0,
1872
1873 fn initBits(comptime T: type, out: anytype, num: u16) Bits(T) {
1874 const UT = @Int(.unsigned, @bitSizeOf(T));
1875 return .{
1876 @bitCast(@as(UT, @intCast(out))),
1877 num,
1878 };
1879 }
1880
1881 fn readBitsNoEof(self: *@This(), comptime T: type, num: u16) !T {
1882 const b, const c = try self.readBitsTuple(T, num);
1883 if (c < num) return error.EndOfStream;
1884 return b;
1885 }
1886
1887 fn readBits(self: *@This(), comptime T: type, num: u16, out_bits: *u16) !T {
1888 const b, const c = try self.readBitsTuple(T, num);
1889 out_bits.* = c;
1890 return b;
1891 }
1892
1893 fn readBitsTuple(self: *@This(), comptime T: type, num: u16) !Bits(T) {
1894 const UT = @Int(.unsigned, @bitSizeOf(T));
1895 const U = if (@bitSizeOf(T) < 8) u8 else UT;
1896
1897 if (num <= self.count) return initBits(T, self.removeBits(@intCast(num)), num);
1898
1899 var out_count: u16 = self.count;
1900 var out: U = self.removeBits(self.count);
1901
1902 const full_bytes_left = (num - out_count) / 8;
1903
1904 for (0..full_bytes_left) |_| {
1905 const byte = takeByte(self) catch |err| switch (err) {
1906 error.EndOfStream => return initBits(T, out, out_count),
1907 };
1908
1909 const pos = @as(U, byte) << @intCast(out_count);
1910 out |= pos;
1911 out_count += 8;
1912 }
1913
1914 const bits_left = num - out_count;
1915 const keep = 8 - bits_left;
1916
1917 if (bits_left == 0) return initBits(T, out, out_count);
1918
1919 const final_byte = takeByte(self) catch |err| switch (err) {
1920 error.EndOfStream => return initBits(T, out, out_count),
1921 };
1922
1923 const pos = @as(U, final_byte & low_bit_mask[bits_left]) << @intCast(out_count);
1924 out |= pos;
1925 self.bits = final_byte >> @intCast(bits_left);
1926
1927 self.count = @intCast(keep);
1928 return initBits(T, out, num);
1929 }
1930
1931 fn takeByte(br: *BitReader) error{EndOfStream}!u8 {
1932 if (br.bytes.len - br.index == 0) return error.EndOfStream;
1933 const result = br.bytes[br.index];
1934 br.index += 1;
1935 return result;
1936 }
1937
1938 fn removeBits(self: *@This(), num: u4) u8 {
1939 if (num == 8) {
1940 self.count = 0;
1941 return self.bits;
1942 }
1943
1944 const keep = self.count - num;
1945 const bits = self.bits & low_bit_mask[num];
1946 self.bits >>= @intCast(num);
1947 self.count = keep;
1948 return bits;
1949 }
1950
1951 fn alignToByte(self: *@This()) void {
1952 self.bits = 0;
1953 self.count = 0;
1954 }
1955};
1956
1957test {
1958 _ = Table;
1959}