authorgravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2025-05-01 14:52:41-07:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2025-07-01 16:35:28-07:00
loge05af2da131b5d9353711ffb5979c67f4bd8b5af
tree11d98d1777956110f2611635181cce9a48fe5ecd
parent685d55c1a49615b6601c02c174d1097489b0fe95

std.compress.zstd: it's compiling


5 files changed, 378 insertions(+), 434 deletions(-)

lib/std/compress/zstd.zig+1-1
...@@ -109,7 +109,7 @@ fn testExpectDecompressError(err: anyerror, compressed: []const u8) !void {...@@ -109,7 +109,7 @@ fn testExpectDecompressError(err: anyerror, compressed: []const u8) !void {
109 in.initFixed(@constCast(compressed));109 in.initFixed(@constCast(compressed));
110 var zstd_stream: Decompress = .init(&in, .{});110 var zstd_stream: Decompress = .init(&in, .{});
111 try std.testing.expectError(error.ReadFailed, zstd_stream.reader().readRemainingArrayList(gpa, null, &out, .unlimited));111 try std.testing.expectError(error.ReadFailed, zstd_stream.reader().readRemainingArrayList(gpa, null, &out, .unlimited));
112 try std.testing.expectError(err, zstd_stream.err.?);112 try std.testing.expectError(err, zstd_stream.err orelse {});
113113
114 return error.TestFailed;114 return error.TestFailed;
115}115}
lib/std/compress/zstd/Decompress.zig+374-133
...@@ -11,12 +11,10 @@ state: State,...@@ -11,12 +11,10 @@ state: State,
11verify_checksum: bool,11verify_checksum: bool,
12err: ?Error = null,12err: ?Error = null,
1313
14const table_size_max = zstd.compressed_block.table_size_max;
15
16const State = union(enum) {14const State = union(enum) {
17 new_frame,15 new_frame,
18 in_frame: InFrame,16 in_frame: InFrame,
19 skipping_frame: u32,17 skipping_frame: usize,
20 end,18 end,
2119
22 const InFrame = struct {20 const InFrame = struct {
...@@ -31,11 +29,38 @@ pub const Options = struct {...@@ -31,11 +29,38 @@ pub const Options = struct {
31};29};
3230
33pub const Error = error{31pub const Error = error{
32 BadMagic,
33 BlockOversize,
34 ChecksumFailure,34 ChecksumFailure,
35 ContentOversize,
35 DictionaryIdFlagUnsupported,36 DictionaryIdFlagUnsupported,
37 EndOfStream,
38 HuffmanTreeIncomplete,
39 InvalidBitStream,
40 LiteralsBufferUndersize,
41 MalformedAccuracyLog,
36 MalformedBlock,42 MalformedBlock,
43 MalformedCompressedBlock,
37 MalformedFrame,44 MalformedFrame,
38 EndOfStream,45 MalformedFseBits,
46 MalformedFseTable,
47 MalformedHuffmanTree,
48 MalformedLiteralsHeader,
49 MalformedLiteralsLength,
50 MalformedLiteralsSection,
51 MalformedSequence,
52 MissingStartBit,
53 OutputBufferUndersize,
54 InputBufferUndersize,
55 ReadFailed,
56 RepeatModeFirst,
57 ReservedBitSet,
58 ReservedBlock,
59 SequenceBufferUndersize,
60 TreelessLiteralsFirst,
61 UnexpectedEndOfLiteralStream,
62 WindowOversize,
63 WindowSizeUnknown,
39};64};
4065
41pub fn init(input: *BufferedReader, options: Options) Decompress {66pub fn init(input: *BufferedReader, options: Options) Decompress {
...@@ -69,25 +94,32 @@ fn read(context: ?*anyopaque, bw: *BufferedWriter, limit: Reader.Limit) Reader.R...@@ -69,25 +94,32 @@ fn read(context: ?*anyopaque, bw: *BufferedWriter, limit: Reader.Limit) Reader.R
69 d.err = err;94 d.err = err;
70 return error.ReadFailed;95 return error.ReadFailed;
71 };96 };
72 return readInFrame(d, bw, limit, &d.state.in_frame) catch |err| {97 return readInFrame(d, bw, limit, &d.state.in_frame) catch |err| switch (err) {
73 d.err = err;98 error.ReadFailed => return error.ReadFailed,
74 return error.ReadFailed;99 error.WriteFailed => return error.WriteFailed,
100 else => |e| {
101 d.err = e;
102 return error.ReadFailed;
103 },
75 };104 };
76 },105 },
77 .in_frame => |*in_frame| {106 .in_frame => |*in_frame| {
78 return readInFrame(d, bw, limit, in_frame) catch |err| {107 return readInFrame(d, bw, limit, in_frame) catch |err| switch (err) {
79 d.err = err;108 error.ReadFailed => return error.ReadFailed,
80 return error.ReadFailed;109 error.WriteFailed => return error.WriteFailed,
110 else => |e| {
111 d.err = e;
112 return error.ReadFailed;
113 },
81 };114 };
82 },115 },
83 .skipping_frame => |*remaining| {116 .skipping_frame => |*remaining| {
84 const requested = remaining.*;117 const n = in.discard(.limited(remaining.*)) catch |err| {
85 const n = in.discard(.limited(requested)) catch |err| {
86 d.err = err;118 d.err = err;
87 return error.ReadFailed;119 return error.ReadFailed;
88 };120 };
89 if (requested == n) d.state = .new_frame;121 remaining.* -= n;
90 remaining.* = requested - n;122 if (remaining.* == 0) d.state = .new_frame;
91 return 0;123 return 0;
92 },124 },
93 .end => return error.EndOfStream,125 .end => return error.EndOfStream,
...@@ -115,9 +147,9 @@ fn initFrame(d: *Decompress, window_size_max: usize, magic: Frame.Magic) !void {...@@ -115,9 +147,9 @@ fn initFrame(d: *Decompress, window_size_max: usize, magic: Frame.Magic) !void {
115fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state: *State.InFrame) !usize {147fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state: *State.InFrame) !usize {
116 const in = d.input;148 const in = d.input;
117149
118 var literal_fse_buffer: [table_size_max.literal]Table.Fse = undefined;150 var literal_fse_buffer: [zstd.table_size_max.literal]Table.Fse = undefined;
119 var match_fse_buffer: [table_size_max.match]Table.Fse = undefined;151 var match_fse_buffer: [zstd.table_size_max.match]Table.Fse = undefined;
120 var offset_fse_buffer: [table_size_max.offset]Table.Fse = undefined;152 var offset_fse_buffer: [zstd.table_size_max.offset]Table.Fse = undefined;
121 var literals_buffer: [zstd.block_size_max]u8 = undefined;153 var literals_buffer: [zstd.block_size_max]u8 = undefined;
122 var sequence_buffer: [zstd.block_size_max]u8 = undefined;154 var sequence_buffer: [zstd.block_size_max]u8 = undefined;
123155
...@@ -125,10 +157,10 @@ fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state:...@@ -125,10 +157,10 @@ fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state:
125157
126 const header_bytes = try in.takeArray(3);158 const header_bytes = try in.takeArray(3);
127 const block_header: Frame.Zstandard.Block.Header = @bitCast(header_bytes.*);159 const block_header: Frame.Zstandard.Block.Header = @bitCast(header_bytes.*);
128 const block_size = block_header.block_size;160 const block_size = block_header.size;
129 if (state.frame.block_size_max < block_size) return error.BlockOversize;161 if (state.frame.block_size_max < block_size) return error.BlockOversize;
130 if (@intFromEnum(limit) < block_size) return error.OutputBufferUndersize;162 if (@intFromEnum(limit) < block_size) return error.OutputBufferUndersize;
131 switch (block_header.block_type) {163 switch (block_header.type) {
132 .raw => {164 .raw => {
133 try in.readAll(bw, .limited(block_size));165 try in.readAll(bw, .limited(block_size));
134 return block_size;166 return block_size;
...@@ -151,9 +183,10 @@ fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state:...@@ -151,9 +183,10 @@ fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state:
151 var bytes_written: usize = 0;183 var bytes_written: usize = 0;
152 {184 {
153 if (sequence_buffer.len < @intFromEnum(remaining))185 if (sequence_buffer.len < @intFromEnum(remaining))
154 return error.SequenceBufferTooSmall;186 return error.SequenceBufferUndersize;
155 const seq_len = try in.readSlice(remaining.slice(&sequence_buffer));187 const seq_slice = remaining.slice(&sequence_buffer);
156 var bit_stream = try ReverseBitReader.init(sequence_buffer[0..seq_len]);188 try in.readSlice(seq_slice);
189 var bit_stream = try ReverseBitReader.init(seq_slice);
157190
158 if (sequences_header.sequence_count > 0) {191 if (sequences_header.sequence_count > 0) {
159 try decode.readInitialFseState(&bit_stream);192 try decode.readInitialFseState(&bit_stream);
...@@ -205,16 +238,16 @@ fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state:...@@ -205,16 +238,16 @@ fn readInFrame(d: *Decompress, bw: *BufferedWriter, limit: Reader.Limit, state:
205 }238 }
206 }239 }
207240
208 if (block_header.last_block) {241 if (block_header.last) {
209 if (state.frame.has_checksum) {242 if (state.frame.has_checksum) {
210 const expected_checksum = try in.readInt(u32, .little);243 const expected_checksum = try in.takeInt(u32, .little);
211 if (state.frame.hasher_opt) |*hasher| {244 if (state.frame.hasher_opt) |*hasher| {
212 const actual_checksum: u32 = @truncate(hasher.final());245 const actual_checksum: u32 = @truncate(hasher.final());
213 if (expected_checksum != actual_checksum) return error.ChecksumFailure;246 if (expected_checksum != actual_checksum) return error.ChecksumFailure;
214 }247 }
215 }248 }
216 if (d.frame.content_size) |content_size| {249 if (state.frame.content_size) |content_size| {
217 if (content_size != d.current_frame_decompressed_size) {250 if (content_size != state.decompressed_size) {
218 return error.MalformedFrame;251 return error.MalformedFrame;
219 }252 }
220 }253 }
...@@ -249,16 +282,16 @@ pub const Frame = struct {...@@ -249,16 +282,16 @@ pub const Frame = struct {
249 _,282 _,
250283
251 pub fn kind(m: Magic) ?Kind {284 pub fn kind(m: Magic) ?Kind {
252 return switch (m) {285 return switch (@intFromEnum(m)) {
253 .zstandard => .zstandard,286 @intFromEnum(Magic.zstandard) => .zstandard,
254 Skippable.magic_min...Skippable.magic_max => .skippable,287 @intFromEnum(Skippable.magic_min)...@intFromEnum(Skippable.magic_max) => .skippable,
255 else => null,288 else => null,
256 };289 };
257 }290 }
258291
259 pub fn isSkippable(m: Magic) bool {292 pub fn isSkippable(m: Magic) bool {
260 return switch (m) {293 return switch (@intFromEnum(m)) {
261 Skippable.magic_min...Skippable.magic_max => true,294 @intFromEnum(Skippable.magic_min)...@intFromEnum(Skippable.magic_max) => true,
262 else => false,295 else => false,
263 };296 };
264 }297 }
...@@ -384,9 +417,9 @@ pub const Frame = struct {...@@ -384,9 +417,9 @@ pub const Frame = struct {
384 ) Decode {417 ) Decode {
385 return .{418 return .{
386 .repeat_offsets = .{419 .repeat_offsets = .{
387 zstd.compressed_block.start_repeated_offset_1,420 zstd.start_repeated_offset_1,
388 zstd.compressed_block.start_repeated_offset_2,421 zstd.start_repeated_offset_2,
389 zstd.compressed_block.start_repeated_offset_3,422 zstd.start_repeated_offset_3,
390 },423 },
391424
392 .offset = undefined,425 .offset = undefined,
...@@ -410,7 +443,7 @@ pub const Frame = struct {...@@ -410,7 +443,7 @@ pub const Frame = struct {
410443
411 pub const PrepareError = error{444 pub const PrepareError = error{
412 /// the (reversed) literal bitstream's first byte does not have any bits set445 /// the (reversed) literal bitstream's first byte does not have any bits set
413 BitStreamHasNoStartBit,446 MissingStartBit,
414 /// `literals` is a treeless literals section and the decode state does not447 /// `literals` is a treeless literals section and the decode state does not
415 /// have a Huffman tree from a previous block448 /// have a Huffman tree from a previous block
416 TreelessLiteralsFirst,449 TreelessLiteralsFirst,
...@@ -422,6 +455,8 @@ pub const Frame = struct {...@@ -422,6 +455,8 @@ pub const Frame = struct {
422 MalformedFseTable,455 MalformedFseTable,
423 /// input stream ends before all FSE tables are read456 /// input stream ends before all FSE tables are read
424 EndOfStream,457 EndOfStream,
458 ReadFailed,
459 InputBufferUndersize,
425 };460 };
426461
427 /// Prepare the decoder to decode a compressed block. Loads the literals462 /// Prepare the decoder to decode a compressed block. Loads the literals
...@@ -430,6 +465,7 @@ pub const Frame = struct {...@@ -430,6 +465,7 @@ pub const Frame = struct {
430 pub fn prepare(465 pub fn prepare(
431 self: *Decode,466 self: *Decode,
432 in: *BufferedReader,467 in: *BufferedReader,
468 remaining: *Reader.Limit,
433 literals: LiteralsSection,469 literals: LiteralsSection,
434 sequences_header: SequencesSection.Header,470 sequences_header: SequencesSection.Header,
435 ) PrepareError!void {471 ) PrepareError!void {
...@@ -455,17 +491,14 @@ pub const Frame = struct {...@@ -455,17 +491,14 @@ pub const Frame = struct {
455 }491 }
456492
457 if (sequences_header.sequence_count > 0) {493 if (sequences_header.sequence_count > 0) {
458 try self.updateFseTable(in, .literal, sequences_header.literal_lengths);494 try self.updateFseTable(in, remaining, .literal, sequences_header.literal_lengths);
459 try self.updateFseTable(in, .offset, sequences_header.offsets);495 try self.updateFseTable(in, remaining, .offset, sequences_header.offsets);
460 try self.updateFseTable(in, .match, sequences_header.match_lengths);496 try self.updateFseTable(in, remaining, .match, sequences_header.match_lengths);
461 self.fse_tables_undefined = false;497 self.fse_tables_undefined = false;
462 }498 }
463 }499 }
464500
465 /// Read initial FSE states for sequence decoding.501 /// Read initial FSE states for sequence decoding.
466 ///
467 /// Errors returned:
468 /// - `error.EndOfStream` if `bit_reader` does not contain enough bits.
469 pub fn readInitialFseState(self: *Decode, bit_reader: *ReverseBitReader) error{EndOfStream}!void {502 pub fn readInitialFseState(self: *Decode, bit_reader: *ReverseBitReader) error{EndOfStream}!void {
470 self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log);503 self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log);
471 self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log);504 self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log);
...@@ -490,6 +523,7 @@ pub const Frame = struct {...@@ -490,6 +523,7 @@ pub const Frame = struct {
490523
491 const DataType = enum { offset, match, literal };524 const DataType = enum { offset, match, literal };
492525
526 /// TODO: don't use `@field`
493 fn updateState(527 fn updateState(
494 self: *Decode,528 self: *Decode,
495 comptime choice: DataType,529 comptime choice: DataType,
...@@ -517,9 +551,11 @@ pub const Frame = struct {...@@ -517,9 +551,11 @@ pub const Frame = struct {
517 EndOfStream,551 EndOfStream,
518 };552 };
519553
554 /// TODO: don't use `@field`
520 fn updateFseTable(555 fn updateFseTable(
521 self: *Decode,556 self: *Decode,
522 source: *BufferedReader,557 in: *BufferedReader,
558 remaining: *Reader.Limit,
523 comptime choice: DataType,559 comptime choice: DataType,
524 mode: SequencesSection.Header.Mode,560 mode: SequencesSection.Header.Mode,
525 ) !void {561 ) !void {
...@@ -527,28 +563,32 @@ pub const Frame = struct {...@@ -527,28 +563,32 @@ pub const Frame = struct {
527 switch (mode) {563 switch (mode) {
528 .predefined => {564 .predefined => {
529 @field(self, field_name).accuracy_log =565 @field(self, field_name).accuracy_log =
530 @field(zstd.compressed_block.default_accuracy_log, field_name);566 @field(zstd.default_accuracy_log, field_name);
531567
532 @field(self, field_name).table =568 @field(self, field_name).table =
533 @field(Table, "predefined_" ++ field_name);569 @field(Table, "predefined_" ++ field_name);
534 },570 },
535 .rle => {571 .rle => {
536 @field(self, field_name).accuracy_log = 0;572 @field(self, field_name).accuracy_log = 0;
537 @field(self, field_name).table = .{ .rle = try source.readByte() };573 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
574 @field(self, field_name).table = .{ .rle = try in.takeByte() };
538 },575 },
539 .fse => {576 .fse => {
540 var bit_reader: std.io.BitReader(.little) = .init(source);577 if (in.buffer.len < @intFromEnum(remaining.*)) return error.InputBufferUndersize;
541578 const limited_buffer = try in.peek(@intFromEnum(remaining.*));
579 var bit_reader: BitReader = .{ .bytes = limited_buffer };
542 const table_size = try Table.decode(580 const table_size = try Table.decode(
543 &bit_reader,581 &bit_reader,
544 @field(zstd.compressed_block.table_symbol_count_max, field_name),582 @field(zstd.table_symbol_count_max, field_name),
545 @field(zstd.compressed_block.table_accuracy_log_max, field_name),583 @field(zstd.table_accuracy_log_max, field_name),
546 @field(self, field_name ++ "_fse_buffer"),584 @field(self, field_name ++ "_fse_buffer"),
547 );585 );
548 @field(self, field_name).table = .{586 @field(self, field_name).table = .{
549 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],587 .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size],
550 };588 };
551 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);589 @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size);
590 in.toss(bit_reader.index);
591 remaining.* = remaining.subtract(bit_reader.index).?;
552 },592 },
553 .repeat => if (self.fse_tables_undefined) return error.RepeatModeFirst,593 .repeat => if (self.fse_tables_undefined) return error.RepeatModeFirst,
554 }594 }
...@@ -571,15 +611,15 @@ pub const Frame = struct {...@@ -571,15 +611,15 @@ pub const Frame = struct {
571 const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code);611 const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code);
572612
573 const match_code = self.getCode(.match);613 const match_code = self.getCode(.match);
574 if (match_code >= zstd.compressed_block.match_length_code_table.len)614 if (match_code >= zstd.match_length_code_table.len)
575 return error.InvalidBitStream;615 return error.InvalidBitStream;
576 const match = zstd.compressed_block.match_length_code_table[match_code];616 const match = zstd.match_length_code_table[match_code];
577 const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]);617 const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]);
578618
579 const literal_code = self.getCode(.literal);619 const literal_code = self.getCode(.literal);
580 if (literal_code >= zstd.compressed_block.literals_length_code_table.len)620 if (literal_code >= zstd.literals_length_code_table.len)
581 return error.InvalidBitStream;621 return error.InvalidBitStream;
582 const literal = zstd.compressed_block.literals_length_code_table[literal_code];622 const literal = zstd.literals_length_code_table[literal_code];
583 const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]);623 const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]);
584624
585 const offset = if (offset_value > 3) offset: {625 const offset = if (offset_value > 3) offset: {
...@@ -622,12 +662,17 @@ pub const Frame = struct {...@@ -622,12 +662,17 @@ pub const Frame = struct {
622 /// The `BufferedWriter` storage capacity is not large enough to662 /// The `BufferedWriter` storage capacity is not large enough to
623 /// accept this stream.663 /// accept this stream.
624 OutputBufferUndersize,664 OutputBufferUndersize,
665 WriteFailed,
666 MalformedLiteralsLength,
667 MalformedFseBits,
668 MissingStartBit,
669 HuffmanTreeIncomplete,
625 };670 };
626671
627 /// Decode one sequence from `bit_reader` into `dest`. Updates FSE states672 /// Decode one sequence from `bit_reader` into `dest`. Updates FSE states
628 /// if `last_sequence` is `false`. Assumes `prepare` called for the block673 /// if `last_sequence` is `false`. Assumes `prepare` called for the block
629 /// before attempting to decode sequences.674 /// before attempting to decode sequences.
630 pub fn decodeSequence(675 fn decodeSequence(
631 self: *Decode,676 self: *Decode,
632 dest: *BufferedWriter,677 dest: *BufferedWriter,
633 bit_reader: *ReverseBitReader,678 bit_reader: *ReverseBitReader,
...@@ -662,13 +707,13 @@ pub const Frame = struct {...@@ -662,13 +707,13 @@ pub const Frame = struct {
662 return sequence_length;707 return sequence_length;
663 }708 }
664709
665 fn nextLiteralMultiStream(self: *Decode) error{BitStreamHasNoStartBit}!void {710 fn nextLiteralMultiStream(self: *Decode) error{MissingStartBit}!void {
666 self.literal_stream_index += 1;711 self.literal_stream_index += 1;
667 try self.initLiteralStream(self.literal_streams.four[self.literal_stream_index]);712 try self.initLiteralStream(self.literal_streams.four[self.literal_stream_index]);
668 }713 }
669714
670 fn initLiteralStream(self: *Decode, bytes: []const u8) error{BitStreamHasNoStartBit}!void {715 fn initLiteralStream(self: *Decode, bytes: []const u8) error{MissingStartBit}!void {
671 try self.literal_stream_reader.init(bytes);716 self.literal_stream_reader = try ReverseBitReader.init(bytes);
672 }717 }
673718
674 fn isLiteralStreamEmpty(self: *Decode) bool {719 fn isLiteralStreamEmpty(self: *Decode) bool {
...@@ -679,7 +724,7 @@ pub const Frame = struct {...@@ -679,7 +724,7 @@ pub const Frame = struct {
679 }724 }
680725
681 const LiteralBitsError = error{726 const LiteralBitsError = error{
682 BitStreamHasNoStartBit,727 MissingStartBit,
683 UnexpectedEndOfLiteralStream,728 UnexpectedEndOfLiteralStream,
684 };729 };
685 fn readLiteralsBits(730 fn readLiteralsBits(
...@@ -704,6 +749,9 @@ pub const Frame = struct {...@@ -704,6 +749,9 @@ pub const Frame = struct {
704 /// Problems decoding Huffman compressed literals749 /// Problems decoding Huffman compressed literals
705 UnexpectedEndOfLiteralStream,750 UnexpectedEndOfLiteralStream,
706 OutputBufferUndersize,751 OutputBufferUndersize,
752 WriteFailed,
753 MissingStartBit,
754 HuffmanTreeIncomplete,
707 };755 };
708756
709 /// Decode `len` bytes of literals into `dest`.757 /// Decode `len` bytes of literals into `dest`.
...@@ -765,6 +813,7 @@ pub const Frame = struct {...@@ -765,6 +813,7 @@ pub const Frame = struct {
765 }813 }
766 }814 }
767815
816 /// TODO: don't use `@field`
768 fn getCode(self: *Decode, comptime choice: DataType) u32 {817 fn getCode(self: *Decode, comptime choice: DataType) u32 {
769 return switch (@field(self, @tagName(choice)).table) {818 return switch (@field(self, @tagName(choice)).table) {
770 .rle => |value| value,819 .rle => |value| value,
...@@ -785,21 +834,17 @@ pub const Frame = struct {...@@ -785,21 +834,17 @@ pub const Frame = struct {
785 };834 };
786835
787 const InitError = error{836 const InitError = error{
837 /// Frame uses a dictionary.
788 DictionaryIdFlagUnsupported,838 DictionaryIdFlagUnsupported,
839 /// Frame does not have a valid window size.
789 WindowSizeUnknown,840 WindowSizeUnknown,
790 WindowTooLarge,841 /// Window size exceeds `window_size_max` or max `usize` value.
791 ContentSizeTooLarge,842 WindowOversize,
843 /// Frame header indicates a content size exceeding max `usize` value.
844 ContentOversize,
792 };845 };
846
793 /// Validates `frame_header` and returns the associated `Frame`.847 /// Validates `frame_header` and returns the associated `Frame`.
794 ///
795 /// Errors returned:
796 /// - `error.DictionaryIdFlagUnsupported` if the frame uses a dictionary
797 /// - `error.WindowSizeUnknown` if the frame does not have a valid window
798 /// size
799 /// - `error.WindowTooLarge` if the window size is larger than
800 /// `window_size_max` or `std.math.intMax(usize)`
801 /// - `error.ContentSizeTooLarge` if the frame header indicates a content
802 /// size larger than `std.math.maxInt(usize)`
803 pub fn init(848 pub fn init(
804 frame_header: Frame.Zstandard.Header,849 frame_header: Frame.Zstandard.Header,
805 window_size_max: usize,850 window_size_max: usize,
...@@ -810,15 +855,15 @@ pub const Frame = struct {...@@ -810,15 +855,15 @@ pub const Frame = struct {
810855
811 const window_size_raw = frame_header.windowSize() orelse return error.WindowSizeUnknown;856 const window_size_raw = frame_header.windowSize() orelse return error.WindowSizeUnknown;
812 const window_size = if (window_size_raw > window_size_max)857 const window_size = if (window_size_raw > window_size_max)
813 return error.WindowTooLarge858 return error.WindowOversize
814 else859 else
815 std.math.cast(usize, window_size_raw) orelse return error.WindowTooLarge;860 std.math.cast(usize, window_size_raw) orelse return error.WindowOversize;
816861
817 const should_compute_checksum =862 const should_compute_checksum =
818 frame_header.descriptor.content_checksum_flag and verify_checksum;863 frame_header.descriptor.content_checksum_flag and verify_checksum;
819864
820 const content_size = if (frame_header.content_size) |size|865 const content_size = if (frame_header.content_size) |size|
821 std.math.cast(usize, size) orelse return error.ContentSizeTooLarge866 std.math.cast(usize, size) orelse return error.ContentOversize
822 else867 else
823 null;868 null;
824869
...@@ -875,13 +920,11 @@ pub const LiteralsSection = struct {...@@ -875,13 +920,11 @@ pub const LiteralsSection = struct {
875 compressed_size: ?u18,920 compressed_size: ?u18,
876921
877 /// Decode a literals section header.922 /// Decode a literals section header.
878 ///923 pub fn decode(in: *BufferedReader, remaining: *Reader.Limit) !Header {
879 /// Errors returned:924 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
880 /// - `error.EndOfStream` if there are not enough bytes in `source`925 const byte0 = try in.takeByte();
881 pub fn decode(source: *BufferedReader) !Header {926 const block_type: BlockType = @enumFromInt(byte0 & 0b11);
882 const byte0 = try source.readByte();927 const size_format: u2 = @intCast((byte0 & 0b1100) >> 2);
883 const block_type = @as(BlockType, @enumFromInt(byte0 & 0b11));
884 const size_format = @as(u2, @intCast((byte0 & 0b1100) >> 2));
885 var regenerated_size: u20 = undefined;928 var regenerated_size: u20 = undefined;
886 var compressed_size: ?u18 = null;929 var compressed_size: ?u18 = null;
887 switch (block_type) {930 switch (block_type) {
...@@ -890,28 +933,37 @@ pub const LiteralsSection = struct {...@@ -890,28 +933,37 @@ pub const LiteralsSection = struct {
890 0, 2 => {933 0, 2 => {
891 regenerated_size = byte0 >> 3;934 regenerated_size = byte0 >> 3;
892 },935 },
893 1 => regenerated_size = (byte0 >> 4) + (@as(u20, try source.readByte()) << 4),936 1 => {
894 3 => regenerated_size = (byte0 >> 4) +937 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
895 (@as(u20, try source.readByte()) << 4) +938 regenerated_size = (byte0 >> 4) + (@as(u20, try in.takeByte()) << 4);
896 (@as(u20, try source.readByte()) << 12),939 },
940 3 => {
941 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
942 regenerated_size = (byte0 >> 4) +
943 (@as(u20, try in.takeByte()) << 4) +
944 (@as(u20, try in.takeByte()) << 12);
945 },
897 }946 }
898 },947 },
899 .compressed, .treeless => {948 .compressed, .treeless => {
900 const byte1 = try source.readByte();949 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
901 const byte2 = try source.readByte();950 const byte1 = try in.takeByte();
951 const byte2 = try in.takeByte();
902 switch (size_format) {952 switch (size_format) {
903 0, 1 => {953 0, 1 => {
904 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);954 regenerated_size = (byte0 >> 4) + ((@as(u20, byte1) & 0b00111111) << 4);
905 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);955 compressed_size = ((byte1 & 0b11000000) >> 6) + (@as(u18, byte2) << 2);
906 },956 },
907 2 => {957 2 => {
908 const byte3 = try source.readByte();958 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
959 const byte3 = try in.takeByte();
909 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);960 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00000011) << 12);
910 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);961 compressed_size = ((byte2 & 0b11111100) >> 2) + (@as(u18, byte3) << 6);
911 },962 },
912 3 => {963 3 => {
913 const byte3 = try source.readByte();964 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
914 const byte4 = try source.readByte();965 const byte3 = try in.takeByte();
966 const byte4 = try in.takeByte();
915 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);967 regenerated_size = (byte0 >> 4) + (@as(u20, byte1) << 4) + ((@as(u20, byte2) & 0b00111111) << 12);
916 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);968 compressed_size = ((byte2 & 0b11000000) >> 6) + (@as(u18, byte3) << 2) + (@as(u18, byte4) << 10);
917 },969 },
...@@ -950,17 +1002,17 @@ pub const LiteralsSection = struct {...@@ -950,17 +1002,17 @@ pub const LiteralsSection = struct {
950 index: usize,1002 index: usize,
951 };1003 };
9521004
953 pub fn query(self: HuffmanTree, index: usize, prefix: u16) error{NotFound}!Result {1005 pub fn query(self: HuffmanTree, index: usize, prefix: u16) error{HuffmanTreeIncomplete}!Result {
954 var node = self.nodes[index];1006 var node = self.nodes[index];
955 const weight = node.weight;1007 const weight = node.weight;
956 var i: usize = index;1008 var i: usize = index;
957 while (node.weight == weight) {1009 while (node.weight == weight) {
958 if (node.prefix == prefix) return Result{ .symbol = node.symbol };1010 if (node.prefix == prefix) return .{ .symbol = node.symbol };
959 if (i == 0) return error.NotFound;1011 if (i == 0) return error.HuffmanTreeIncomplete;
960 i -= 1;1012 i -= 1;
961 node = self.nodes[i];1013 node = self.nodes[i];
962 }1014 }
963 return Result{ .index = i };1015 return .{ .index = i };
964 }1016 }
9651017
966 pub fn weightToBitCount(weight: u4, max_bit_count: u4) u4 {1018 pub fn weightToBitCount(weight: u4, max_bit_count: u4) u4 {
...@@ -975,20 +1027,26 @@ pub const LiteralsSection = struct {...@@ -975,20 +1027,26 @@ pub const LiteralsSection = struct {
975 MissingStartBit,1027 MissingStartBit,
976 };1028 };
9771029
978 pub fn decode(in: *BufferedReader) HuffmanTree.DecodeError!HuffmanTree {1030 pub fn decode(in: *BufferedReader, remaining: *Reader.Limit) HuffmanTree.DecodeError!HuffmanTree {
1031 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
979 const header = try in.takeByte();1032 const header = try in.takeByte();
980 if (header < 128) {1033 if (header < 128) {
981 return decodeFse(in, header);1034 return decodeFse(in, remaining, header);
982 } else {1035 } else {
983 return decodeDirect(in, header - 127);1036 return decodeDirect(in, remaining, header - 127);
984 }1037 }
985 }1038 }
9861039
987 fn decodeDirect(source: *BufferedReader, encoded_symbol_count: usize) HuffmanTree.DecodeError!HuffmanTree {1040 fn decodeDirect(
1041 in: *BufferedReader,
1042 remaining: *Reader.Limit,
1043 encoded_symbol_count: usize,
1044 ) HuffmanTree.DecodeError!HuffmanTree {
988 var weights: [256]u4 = undefined;1045 var weights: [256]u4 = undefined;
989 const weights_byte_count = (encoded_symbol_count + 1) / 2;1046 const weights_byte_count = (encoded_symbol_count + 1) / 2;
1047 remaining.* = remaining.subtract(weights_byte_count) orelse return error.EndOfStream;
990 for (0..weights_byte_count) |i| {1048 for (0..weights_byte_count) |i| {
991 const byte = try source.takeByte();1049 const byte = try in.takeByte();
992 weights[2 * i] = @as(u4, @intCast(byte >> 4));1050 weights[2 * i] = @as(u4, @intCast(byte >> 4));
993 weights[2 * i + 1] = @as(u4, @intCast(byte & 0xF));1051 weights[2 * i + 1] = @as(u4, @intCast(byte & 0xF));
994 }1052 }
...@@ -996,22 +1054,25 @@ pub const LiteralsSection = struct {...@@ -996,22 +1054,25 @@ pub const LiteralsSection = struct {
996 return build(&weights, symbol_count);1054 return build(&weights, symbol_count);
997 }1055 }
9981056
999 fn decodeFse(in: *BufferedReader, compressed_size: usize) HuffmanTree.DecodeError!HuffmanTree {1057 fn decodeFse(
1058 in: *BufferedReader,
1059 remaining: *Reader.Limit,
1060 compressed_size: usize,
1061 ) HuffmanTree.DecodeError!HuffmanTree {
1000 var weights: [256]u4 = undefined;1062 var weights: [256]u4 = undefined;
1063 remaining.* = remaining.subtract(compressed_size) orelse return error.EndOfStream;
1001 const compressed_buffer = try in.take(compressed_size);1064 const compressed_buffer = try in.take(compressed_size);
1002 var limited_stream: BufferedReader = undefined;1065 var bit_reader: BitReader = .{ .bytes = compressed_buffer };
1003 limited_stream.initFixed(compressed_buffer);
1004 var bit_reader: std.io.BitReader(.little) = .init(&limited_stream);
1005 var entries: [1 << 6]Table.Fse = undefined;1066 var entries: [1 << 6]Table.Fse = undefined;
1006 const table_size = try Table.decode(&bit_reader, 256, 6, &entries);1067 const table_size = try Table.decode(&bit_reader, 256, 6, &entries);
1007 const accuracy_log = std.math.log2_int_ceil(usize, table_size);1068 const accuracy_log = std.math.log2_int_ceil(usize, table_size);
1008 const remaining = limited_stream.bufferContents();1069 const remaining_buffer = bit_reader.bytes[bit_reader.index..];
1009 const symbol_count = try assignWeights(remaining, accuracy_log, &entries, weights);1070 const symbol_count = try assignWeights(remaining_buffer, accuracy_log, &entries, &weights);
1010 return build(&weights, symbol_count);1071 return build(&weights, symbol_count);
1011 }1072 }
10121073
1013 fn assignWeights(1074 fn assignWeights(
1014 huff_bits_buffer: []u8,1075 huff_bits_buffer: []const u8,
1015 accuracy_log: u16,1076 accuracy_log: u16,
1016 entries: *[1 << 6]Table.Fse,1077 entries: *[1 << 6]Table.Fse,
1017 weights: *[256]u4,1078 weights: *[256]u4,
...@@ -1159,14 +1220,18 @@ pub const LiteralsSection = struct {...@@ -1159,14 +1220,18 @@ pub const LiteralsSection = struct {
1159 MalformedHuffmanTree,1220 MalformedHuffmanTree,
1160 /// Not enough bytes to complete the section.1221 /// Not enough bytes to complete the section.
1161 EndOfStream,1222 EndOfStream,
1223 ReadFailed,
1224 LiteralsBufferUndersize,
1225 MissingStartBit,
1162 };1226 };
11631227
1164 pub fn decode(source: *BufferedReader, buffer: []u8) DecodeError!LiteralsSection {1228 pub fn decode(in: *BufferedReader, remaining: *Reader.Limit, buffer: []u8) DecodeError!LiteralsSection {
1165 const header = try Header.decode(source);1229 const header = try Header.decode(in, remaining);
1166 switch (header.block_type) {1230 switch (header.block_type) {
1167 .raw => {1231 .raw => {
1168 if (buffer.len < header.regenerated_size) return error.LiteralsBufferTooSmall;1232 if (buffer.len < header.regenerated_size) return error.LiteralsBufferUndersize;
1169 try source.readNoEof(buffer[0..header.regenerated_size]);1233 remaining.* = remaining.subtract(header.regenerated_size) orelse return error.EndOfStream;
1234 try in.readSlice(buffer[0..header.regenerated_size]);
1170 return .{1235 return .{
1171 .header = header,1236 .header = header,
1172 .huffman_tree = null,1237 .huffman_tree = null,
...@@ -1174,7 +1239,8 @@ pub const LiteralsSection = struct {...@@ -1174,7 +1239,8 @@ pub const LiteralsSection = struct {
1174 };1239 };
1175 },1240 },
1176 .rle => {1241 .rle => {
1177 buffer[0] = try source.readByte();1242 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1243 buffer[0] = try in.takeByte();
1178 return .{1244 return .{
1179 .header = header,1245 .header = header,
1180 .huffman_tree = null,1246 .huffman_tree = null,
...@@ -1182,19 +1248,18 @@ pub const LiteralsSection = struct {...@@ -1182,19 +1248,18 @@ pub const LiteralsSection = struct {
1182 };1248 };
1183 },1249 },
1184 .compressed, .treeless => {1250 .compressed, .treeless => {
1185 var counting_reader = std.io.countingReader(source);1251 const before_remaining = remaining.*;
1186 const huffman_tree = if (header.block_type == .compressed)1252 const huffman_tree = if (header.block_type == .compressed)
1187 try HuffmanTree.decode(counting_reader.reader(), buffer)1253 try HuffmanTree.decode(in, remaining)
1188 else1254 else
1189 null;1255 null;
1190 const huffman_tree_size = @as(usize, @intCast(counting_reader.bytes_read));1256 const huffman_tree_size = @intFromEnum(before_remaining) - @intFromEnum(remaining.*);
1191 const total_streams_size = std.math.sub(usize, header.compressed_size.?, huffman_tree_size) catch1257 const total_streams_size = std.math.sub(usize, header.compressed_size.?, huffman_tree_size) catch
1192 return error.MalformedLiteralsSection;1258 return error.MalformedLiteralsSection;
11931259 if (total_streams_size > buffer.len) return error.LiteralsBufferUndersize;
1194 if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall;1260 remaining.* = remaining.subtract(total_streams_size) orelse return error.EndOfStream;
1195 try source.readNoEof(buffer[0..total_streams_size]);1261 try in.readSlice(buffer[0..total_streams_size]);
1196 const stream_data = buffer[0..total_streams_size];1262 const stream_data = buffer[0..total_streams_size];
1197
1198 const streams = try Streams.decode(header.size_format, stream_data);1263 const streams = try Streams.decode(header.size_format, stream_data);
1199 return .{1264 return .{
1200 .header = header,1265 .header = header,
...@@ -1207,7 +1272,7 @@ pub const LiteralsSection = struct {...@@ -1207,7 +1272,7 @@ pub const LiteralsSection = struct {
1207};1272};
12081273
1209pub const SequencesSection = struct {1274pub const SequencesSection = struct {
1210 header: SequencesSection.Header,1275 header: Header,
1211 literals_length_table: Table,1276 literals_length_table: Table,
1212 offset_table: Table,1277 offset_table: Table,
1213 match_length_table: Table,1278 match_length_table: Table,
...@@ -1228,32 +1293,37 @@ pub const SequencesSection = struct {...@@ -1228,32 +1293,37 @@ pub const SequencesSection = struct {
1228 pub const DecodeError = error{1293 pub const DecodeError = error{
1229 ReservedBitSet,1294 ReservedBitSet,
1230 EndOfStream,1295 EndOfStream,
1296 ReadFailed,
1231 };1297 };
12321298
1233 pub fn decode(source: *BufferedReader) DecodeError!Header {1299 pub fn decode(in: *BufferedReader, remaining: *Reader.Limit) DecodeError!Header {
1234 var sequence_count: u24 = undefined;1300 var sequence_count: u24 = undefined;
12351301
1236 const byte0 = try source.readByte();1302 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1303 const byte0 = try in.takeByte();
1237 if (byte0 == 0) {1304 if (byte0 == 0) {
1238 return SequencesSection.Header{1305 return .{
1239 .sequence_count = 0,1306 .sequence_count = 0,
1240 .offsets = undefined,1307 .offsets = undefined,
1241 .match_lengths = undefined,1308 .match_lengths = undefined,
1242 .literal_lengths = undefined,1309 .literal_lengths = undefined,
1243 };1310 };
1244 } else if (byte0 < 128) {1311 } else if (byte0 < 128) {
1312 remaining.* = remaining.subtract(1) orelse return error.EndOfStream;
1245 sequence_count = byte0;1313 sequence_count = byte0;
1246 } else if (byte0 < 255) {1314 } else if (byte0 < 255) {
1247 sequence_count = (@as(u24, (byte0 - 128)) << 8) + try source.readByte();1315 remaining.* = remaining.subtract(2) orelse return error.EndOfStream;
1316 sequence_count = (@as(u24, (byte0 - 128)) << 8) + try in.takeByte();
1248 } else {1317 } else {
1249 sequence_count = (try source.readByte()) + (@as(u24, try source.readByte()) << 8) + 0x7F00;1318 remaining.* = remaining.subtract(3) orelse return error.EndOfStream;
1319 sequence_count = (try in.takeByte()) + (@as(u24, try in.takeByte()) << 8) + 0x7F00;
1250 }1320 }
12511321
1252 const compression_modes = try source.readByte();1322 const compression_modes = try in.takeByte();
12531323
1254 const matches_mode = @as(SequencesSection.Header.Mode, @enumFromInt((compression_modes & 0b00001100) >> 2));1324 const matches_mode: Header.Mode = @enumFromInt((compression_modes & 0b00001100) >> 2);
1255 const offsets_mode = @as(SequencesSection.Header.Mode, @enumFromInt((compression_modes & 0b00110000) >> 4));1325 const offsets_mode: Header.Mode = @enumFromInt((compression_modes & 0b00110000) >> 4);
1256 const literal_mode = @as(SequencesSection.Header.Mode, @enumFromInt((compression_modes & 0b11000000) >> 6));1326 const literal_mode: Header.Mode = @enumFromInt((compression_modes & 0b11000000) >> 6);
1257 if (compression_modes & 0b11 != 0) return error.ReservedBitSet;1327 if (compression_modes & 0b11 != 0) return error.ReservedBitSet;
12581328
1259 return .{1329 return .{
...@@ -1277,7 +1347,7 @@ pub const Table = union(enum) {...@@ -1277,7 +1347,7 @@ pub const Table = union(enum) {
1277 };1347 };
12781348
1279 pub fn decode(1349 pub fn decode(
1280 bit_reader: *std.io.BitReader(.little),1350 bit_reader: *BitReader,
1281 expected_symbol_count: usize,1351 expected_symbol_count: usize,
1282 max_accuracy_log: u4,1352 max_accuracy_log: u4,
1283 entries: []Table.Fse,1353 entries: []Table.Fse,
...@@ -1600,6 +1670,22 @@ pub const Table = union(enum) {...@@ -1600,6 +1670,22 @@ pub const Table = union(enum) {
1600 };1670 };
1601};1671};
16021672
1673const low_bit_mask = [9]u8{
1674 0b00000000,
1675 0b00000001,
1676 0b00000011,
1677 0b00000111,
1678 0b00001111,
1679 0b00011111,
1680 0b00111111,
1681 0b01111111,
1682 0b11111111,
1683};
1684
1685fn Bits(comptime T: type) type {
1686 return struct { T, u16 };
1687}
1688
1603/// For reading the reversed bit streams used to encode FSE compressed data.1689/// For reading the reversed bit streams used to encode FSE compressed data.
1604const ReverseBitReader = struct {1690const ReverseBitReader = struct {
1605 bytes: []const u8,1691 bytes: []const u8,
...@@ -1619,20 +1705,175 @@ const ReverseBitReader = struct {...@@ -1619,20 +1705,175 @@ const ReverseBitReader = struct {
1619 return error.MissingStartBit;1705 return error.MissingStartBit;
1620 }1706 }
16211707
1622 fn readBitsNoEof(self: *ReverseBitReader, comptime U: type, num_bits: u16) error{EndOfStream}!U {1708 fn initBits(comptime T: type, out: anytype, num: u16) Bits(T) {
1623 return self.bit_reader.readBitsNoEof(U, num_bits);1709 const UT = std.meta.Int(.unsigned, @bitSizeOf(T));
1710 return .{
1711 @bitCast(@as(UT, @intCast(out))),
1712 num,
1713 };
1714 }
1715
1716 fn readBitsNoEof(self: *ReverseBitReader, comptime T: type, num: u16) error{EndOfStream}!T {
1717 const b, const c = try self.readBitsTuple(T, num);
1718 if (c < num) return error.EndOfStream;
1719 return b;
1624 }1720 }
16251721
1626 fn readBits(self: *ReverseBitReader, comptime U: type, num_bits: u16, out_bits: *u16) error{}!U {1722 fn readBits(self: *ReverseBitReader, comptime T: type, num: u16, out_bits: *u16) !T {
1627 return try self.bit_reader.readBits(U, num_bits, out_bits);1723 const b, const c = try self.readBitsTuple(T, num);
1724 out_bits.* = c;
1725 return b;
1628 }1726 }
16291727
1630 fn alignToByte(self: *ReverseBitReader) void {1728 fn readBitsTuple(self: *ReverseBitReader, comptime T: type, num: u16) !Bits(T) {
1631 self.bit_reader.alignToByte();1729 const UT = std.meta.Int(.unsigned, @bitSizeOf(T));
1730 const U = if (@bitSizeOf(T) < 8) u8 else UT;
1731
1732 if (num <= self.count) return initBits(T, self.removeBits(@intCast(num)), num);
1733
1734 var out_count: u16 = self.count;
1735 var out: U = self.removeBits(self.count);
1736
1737 const full_bytes_left = (num - out_count) / 8;
1738
1739 for (0..full_bytes_left) |_| {
1740 const byte = takeByte(self) catch |err| switch (err) {
1741 error.EndOfStream => return initBits(T, out, out_count),
1742 };
1743 if (U == u8) out = 0 else out <<= 8;
1744 out |= byte;
1745 out_count += 8;
1746 }
1747
1748 const bits_left = num - out_count;
1749 const keep = 8 - bits_left;
1750
1751 if (bits_left == 0) return initBits(T, out, out_count);
1752
1753 const final_byte = takeByte(self) catch |err| switch (err) {
1754 error.EndOfStream => return initBits(T, out, out_count),
1755 };
1756
1757 out <<= @intCast(bits_left);
1758 out |= final_byte >> @intCast(keep);
1759 self.bits = final_byte & low_bit_mask[keep];
1760
1761 self.count = @intCast(keep);
1762 return initBits(T, out, num);
1763 }
1764
1765 fn takeByte(rbr: *ReverseBitReader) error{EndOfStream}!u8 {
1766 if (rbr.remaining == 0) return error.EndOfStream;
1767 rbr.remaining -= 1;
1768 return rbr.bytes[rbr.remaining];
1632 }1769 }
16331770
1634 fn isEmpty(self: *const ReverseBitReader) bool {1771 fn isEmpty(self: *const ReverseBitReader) bool {
1635 return self.byte_reader.remaining_bytes == 0 and self.bit_reader.count == 0;1772 return self.remaining == 0 and self.count == 0;
1773 }
1774
1775 fn removeBits(self: *ReverseBitReader, num: u4) u8 {
1776 if (num == 8) {
1777 self.count = 0;
1778 return self.bits;
1779 }
1780
1781 const keep = self.count - num;
1782 const bits = self.bits >> @intCast(keep);
1783 self.bits &= low_bit_mask[keep];
1784
1785 self.count = keep;
1786 return bits;
1787 }
1788};
1789
1790const BitReader = struct {
1791 bytes: []const u8,
1792 index: usize = 0,
1793 bits: u8 = 0,
1794 count: u4 = 0,
1795
1796 fn initBits(comptime T: type, out: anytype, num: u16) Bits(T) {
1797 const UT = std.meta.Int(.unsigned, @bitSizeOf(T));
1798 return .{
1799 @bitCast(@as(UT, @intCast(out))),
1800 num,
1801 };
1802 }
1803
1804 fn readBitsNoEof(self: *@This(), comptime T: type, num: u16) !T {
1805 const b, const c = try self.readBitsTuple(T, num);
1806 if (c < num) return error.EndOfStream;
1807 return b;
1808 }
1809
1810 fn readBits(self: *@This(), comptime T: type, num: u16, out_bits: *u16) !T {
1811 const b, const c = try self.readBitsTuple(T, num);
1812 out_bits.* = c;
1813 return b;
1814 }
1815
1816 fn readBitsTuple(self: *@This(), comptime T: type, num: u16) !Bits(T) {
1817 const UT = std.meta.Int(.unsigned, @bitSizeOf(T));
1818 const U = if (@bitSizeOf(T) < 8) u8 else UT;
1819
1820 if (num <= self.count) return initBits(T, self.removeBits(@intCast(num)), num);
1821
1822 var out_count: u16 = self.count;
1823 var out: U = self.removeBits(self.count);
1824
1825 const full_bytes_left = (num - out_count) / 8;
1826
1827 for (0..full_bytes_left) |_| {
1828 const byte = takeByte(self) catch |err| switch (err) {
1829 error.EndOfStream => return initBits(T, out, out_count),
1830 };
1831
1832 const pos = @as(U, byte) << @intCast(out_count);
1833 out |= pos;
1834 out_count += 8;
1835 }
1836
1837 const bits_left = num - out_count;
1838 const keep = 8 - bits_left;
1839
1840 if (bits_left == 0) return initBits(T, out, out_count);
1841
1842 const final_byte = takeByte(self) catch |err| switch (err) {
1843 error.EndOfStream => return initBits(T, out, out_count),
1844 };
1845
1846 const pos = @as(U, final_byte & low_bit_mask[bits_left]) << @intCast(out_count);
1847 out |= pos;
1848 self.bits = final_byte >> @intCast(bits_left);
1849
1850 self.count = @intCast(keep);
1851 return initBits(T, out, num);
1852 }
1853
1854 fn takeByte(br: *BitReader) error{EndOfStream}!u8 {
1855 if (br.bytes.len - br.index == 0) return error.EndOfStream;
1856 const result = br.bytes[br.index];
1857 br.index += 1;
1858 return result;
1859 }
1860
1861 fn removeBits(self: *@This(), num: u4) u8 {
1862 if (num == 8) {
1863 self.count = 0;
1864 return self.bits;
1865 }
1866
1867 const keep = self.count - num;
1868 const bits = self.bits & low_bit_mask[num];
1869 self.bits >>= @intCast(num);
1870 self.count = keep;
1871 return bits;
1872 }
1873
1874 fn alignToByte(self: *@This()) void {
1875 self.bits = 0;
1876 self.count = 0;
1636 }1877 }
1637};1878};
16381879
lib/std/io.zig+3-9
...@@ -19,16 +19,12 @@ pub const AllocatingWriter = @import("io/AllocatingWriter.zig");...@@ -19,16 +19,12 @@ pub const AllocatingWriter = @import("io/AllocatingWriter.zig");
19pub const MultiWriter = @import("io/multi_writer.zig").MultiWriter;19pub const MultiWriter = @import("io/multi_writer.zig").MultiWriter;
20pub const multiWriter = @import("io/multi_writer.zig").multiWriter;20pub const multiWriter = @import("io/multi_writer.zig").multiWriter;
2121
22pub const BitReader = @import("io/bit_reader.zig").Type;
23
24pub const BitWriter = @import("io/bit_writer.zig").BitWriter;22pub const BitWriter = @import("io/bit_writer.zig").BitWriter;
25pub const bitWriter = @import("io/bit_writer.zig").bitWriter;23pub const bitWriter = @import("io/bit_writer.zig").bitWriter;
2624
27pub const ChangeDetectionStream = @import("io/change_detection_stream.zig").ChangeDetectionStream;25pub const ChangeDetectionStream = @import("io/change_detection_stream.zig").ChangeDetectionStream;
28pub const changeDetectionStream = @import("io/change_detection_stream.zig").changeDetectionStream;26pub const changeDetectionStream = @import("io/change_detection_stream.zig").changeDetectionStream;
2927
30pub const BufferedAtomicFile = @import("io/buffered_atomic_file.zig").BufferedAtomicFile;
31
32pub const tty = @import("io/tty.zig");28pub const tty = @import("io/tty.zig");
3329
34pub fn poll(30pub fn poll(
...@@ -437,13 +433,11 @@ pub fn PollFiles(comptime StreamEnum: type) type {...@@ -437,13 +433,11 @@ pub fn PollFiles(comptime StreamEnum: type) type {
437}433}
438434
439test {435test {
440 _ = BufferedWriter;436 _ = AllocatingWriter;
437 _ = BitWriter;
441 _ = BufferedReader;438 _ = BufferedReader;
439 _ = BufferedWriter;
442 _ = Reader;440 _ = Reader;
443 _ = Writer;441 _ = Writer;
444 _ = AllocatingWriter;
445 _ = @import("io/bit_reader.zig");
446 _ = @import("io/bit_writer.zig");
447 _ = @import("io/buffered_atomic_file.zig");
448 _ = @import("io/test.zig");442 _ = @import("io/test.zig");
449}443}
lib/std/io/bit_reader.zig deleted-236
...@@ -1,236 +0,0 @@
1const std = @import("../std.zig");
2const bit_reader = @This();
3
4//General note on endianess:
5//Big endian is packed starting in the most significant part of the byte and subsequent
6// bytes contain less significant bits. Thus we always take bits from the high
7// end and place them below existing bits in our output.
8//Little endian is packed starting in the least significant part of the byte and
9// subsequent bytes contain more significant bits. Thus we always take bits from
10// the low end and place them above existing bits in our output.
11//Regardless of endianess, within any given byte the bits are always in most
12// to least significant order.
13//Also regardless of endianess, the buffer always aligns bits to the low end
14// of the byte.
15
16/// Creates a bit reader which allows for reading bits from an underlying standard reader
17pub fn Type(comptime endian: std.builtin.Endian) type {
18 return struct {
19 reader: *std.io.BufferedReader,
20 bits: u8,
21 count: u4,
22
23 const low_bit_mask = [9]u8{
24 0b00000000,
25 0b00000001,
26 0b00000011,
27 0b00000111,
28 0b00001111,
29 0b00011111,
30 0b00111111,
31 0b01111111,
32 0b11111111,
33 };
34
35 pub fn init(reader: *std.io.BufferedReader) @This() {
36 return .{ .reader = reader, .bits = 0, .count = 0 };
37 }
38
39 fn Bits(comptime T: type) type {
40 return struct { T, u16 };
41 }
42
43 fn initBits(comptime T: type, out: anytype, num: u16) Bits(T) {
44 const UT = std.meta.Int(.unsigned, @bitSizeOf(T));
45 return .{
46 @bitCast(@as(UT, @intCast(out))),
47 num,
48 };
49 }
50
51 /// Reads `bits` bits from the reader and returns a specified type
52 /// containing them in the least significant end, returning an error if the
53 /// specified number of bits could not be read.
54 pub fn readBitsNoEof(self: *@This(), comptime T: type, num: u16) !T {
55 const b, const c = try self.readBitsTuple(T, num);
56 if (c < num) return error.EndOfStream;
57 return b;
58 }
59
60 /// Reads `bits` bits from the reader and returns a specified type
61 /// containing them in the least significant end. The number of bits successfully
62 /// read is placed in `out_bits`, as reaching the end of the stream is not an error.
63 pub fn readBits(self: *@This(), comptime T: type, num: u16, out_bits: *u16) !T {
64 const b, const c = try self.readBitsTuple(T, num);
65 out_bits.* = c;
66 return b;
67 }
68
69 /// Reads `bits` bits from the reader and returns a tuple of the specified type
70 /// containing them in the least significant end, and the number of bits successfully
71 /// read. Reaching the end of the stream is not an error.
72 pub fn readBitsTuple(self: *@This(), comptime T: type, num: u16) !Bits(T) {
73 const UT = std.meta.Int(.unsigned, @bitSizeOf(T));
74 const U = if (@bitSizeOf(T) < 8) u8 else UT; //it is a pain to work with <u8
75
76 //dump any bits in our buffer first
77 if (num <= self.count) return initBits(T, self.removeBits(@intCast(num)), num);
78
79 var out_count: u16 = self.count;
80 var out: U = self.removeBits(self.count);
81
82 //grab all the full bytes we need and put their
83 //bits where they belong
84 const full_bytes_left = (num - out_count) / 8;
85
86 for (0..full_bytes_left) |_| {
87 const byte = self.reader.takeByte() catch |err| switch (err) {
88 error.EndOfStream => return initBits(T, out, out_count),
89 else => |e| return e,
90 };
91
92 switch (endian) {
93 .big => {
94 if (U == u8) out = 0 else out <<= 8; //shifting u8 by 8 is illegal in Zig
95 out |= byte;
96 },
97 .little => {
98 const pos = @as(U, byte) << @intCast(out_count);
99 out |= pos;
100 },
101 }
102 out_count += 8;
103 }
104
105 const bits_left = num - out_count;
106 const keep = 8 - bits_left;
107
108 if (bits_left == 0) return initBits(T, out, out_count);
109
110 const final_byte = self.reader.takeByte() catch |err| switch (err) {
111 error.EndOfStream => return initBits(T, out, out_count),
112 else => |e| return e,
113 };
114
115 switch (endian) {
116 .big => {
117 out <<= @intCast(bits_left);
118 out |= final_byte >> @intCast(keep);
119 self.bits = final_byte & low_bit_mask[keep];
120 },
121 .little => {
122 const pos = @as(U, final_byte & low_bit_mask[bits_left]) << @intCast(out_count);
123 out |= pos;
124 self.bits = final_byte >> @intCast(bits_left);
125 },
126 }
127
128 self.count = @intCast(keep);
129 return initBits(T, out, num);
130 }
131
132 //convenience function for removing bits from
133 //the appropriate part of the buffer based on
134 //endianess.
135 fn removeBits(self: *@This(), num: u4) u8 {
136 if (num == 8) {
137 self.count = 0;
138 return self.bits;
139 }
140
141 const keep = self.count - num;
142 const bits = switch (endian) {
143 .big => self.bits >> @intCast(keep),
144 .little => self.bits & low_bit_mask[num],
145 };
146 switch (endian) {
147 .big => self.bits &= low_bit_mask[keep],
148 .little => self.bits >>= @intCast(num),
149 }
150
151 self.count = keep;
152 return bits;
153 }
154
155 pub fn alignToByte(self: *@This()) void {
156 self.bits = 0;
157 self.count = 0;
158 }
159 };
160}
161
162///////////////////////////////
163
164test "api coverage" {
165 const mem_be = [_]u8{ 0b11001101, 0b00001011 };
166 const mem_le = [_]u8{ 0b00011101, 0b10010101 };
167
168 var mem_in_be = std.io.fixedBufferStream(&mem_be);
169 var bit_stream_be: bit_reader.Type(.big) = .init(mem_in_be.reader());
170
171 var out_bits: u16 = undefined;
172
173 const expect = std.testing.expect;
174 const expectError = std.testing.expectError;
175
176 try expect(1 == try bit_stream_be.readBits(u2, 1, &out_bits));
177 try expect(out_bits == 1);
178 try expect(2 == try bit_stream_be.readBits(u5, 2, &out_bits));
179 try expect(out_bits == 2);
180 try expect(3 == try bit_stream_be.readBits(u128, 3, &out_bits));
181 try expect(out_bits == 3);
182 try expect(4 == try bit_stream_be.readBits(u8, 4, &out_bits));
183 try expect(out_bits == 4);
184 try expect(5 == try bit_stream_be.readBits(u9, 5, &out_bits));
185 try expect(out_bits == 5);
186 try expect(1 == try bit_stream_be.readBits(u1, 1, &out_bits));
187 try expect(out_bits == 1);
188
189 mem_in_be.pos = 0;
190 bit_stream_be.count = 0;
191 try expect(0b110011010000101 == try bit_stream_be.readBits(u15, 15, &out_bits));
192 try expect(out_bits == 15);
193
194 mem_in_be.pos = 0;
195 bit_stream_be.count = 0;
196 try expect(0b1100110100001011 == try bit_stream_be.readBits(u16, 16, &out_bits));
197 try expect(out_bits == 16);
198
199 _ = try bit_stream_be.readBits(u0, 0, &out_bits);
200
201 try expect(0 == try bit_stream_be.readBits(u1, 1, &out_bits));
202 try expect(out_bits == 0);
203 try expectError(error.EndOfStream, bit_stream_be.readBitsNoEof(u1, 1));
204
205 var mem_in_le = std.io.fixedBufferStream(&mem_le);
206 var bit_stream_le: bit_reader.Type(.little) = .init(mem_in_le.reader());
207
208 try expect(1 == try bit_stream_le.readBits(u2, 1, &out_bits));
209 try expect(out_bits == 1);
210 try expect(2 == try bit_stream_le.readBits(u5, 2, &out_bits));
211 try expect(out_bits == 2);
212 try expect(3 == try bit_stream_le.readBits(u128, 3, &out_bits));
213 try expect(out_bits == 3);
214 try expect(4 == try bit_stream_le.readBits(u8, 4, &out_bits));
215 try expect(out_bits == 4);
216 try expect(5 == try bit_stream_le.readBits(u9, 5, &out_bits));
217 try expect(out_bits == 5);
218 try expect(1 == try bit_stream_le.readBits(u1, 1, &out_bits));
219 try expect(out_bits == 1);
220
221 mem_in_le.pos = 0;
222 bit_stream_le.count = 0;
223 try expect(0b001010100011101 == try bit_stream_le.readBits(u15, 15, &out_bits));
224 try expect(out_bits == 15);
225
226 mem_in_le.pos = 0;
227 bit_stream_le.count = 0;
228 try expect(0b1001010100011101 == try bit_stream_le.readBits(u16, 16, &out_bits));
229 try expect(out_bits == 16);
230
231 _ = try bit_stream_le.readBits(u0, 0, &out_bits);
232
233 try expect(0 == try bit_stream_le.readBits(u1, 1, &out_bits));
234 try expect(out_bits == 0);
235 try expectError(error.EndOfStream, bit_stream_le.readBitsNoEof(u1, 1));
236}
lib/std/io/buffered_atomic_file.zig deleted-55
...@@ -1,55 +0,0 @@
1const std = @import("../std.zig");
2const mem = std.mem;
3const fs = std.fs;
4const File = std.fs.File;
5
6pub const BufferedAtomicFile = struct {
7 atomic_file: fs.AtomicFile,
8 file_writer: File.Writer,
9 buffered_writer: BufferedWriter,
10 allocator: mem.Allocator,
11
12 pub const buffer_size = 4096;
13 pub const BufferedWriter = std.io.BufferedWriter(buffer_size, File.Writer);
14 pub const Writer = std.io.Writer(*BufferedWriter, BufferedWriter.Error, BufferedWriter.write);
15
16 /// TODO when https://github.com/ziglang/zig/issues/2761 is solved
17 /// this API will not need an allocator
18 pub fn create(
19 allocator: mem.Allocator,
20 dir: fs.Dir,
21 dest_path: []const u8,
22 atomic_file_options: fs.Dir.AtomicFileOptions,
23 ) !*BufferedAtomicFile {
24 var self = try allocator.create(BufferedAtomicFile);
25 self.* = BufferedAtomicFile{
26 .atomic_file = undefined,
27 .file_writer = undefined,
28 .buffered_writer = undefined,
29 .allocator = allocator,
30 };
31 errdefer allocator.destroy(self);
32
33 self.atomic_file = try dir.atomicFile(dest_path, atomic_file_options);
34 errdefer self.atomic_file.deinit();
35
36 self.file_writer = self.atomic_file.file.writer();
37 self.buffered_writer = .{ .unbuffered_writer = self.file_writer };
38 return self;
39 }
40
41 /// always call destroy, even after successful finish()
42 pub fn destroy(self: *BufferedAtomicFile) void {
43 self.atomic_file.deinit();
44 self.allocator.destroy(self);
45 }
46
47 pub fn finish(self: *BufferedAtomicFile) !void {
48 try self.buffered_writer.flush();
49 try self.atomic_file.finish();
50 }
51
52 pub fn writer(self: *BufferedAtomicFile) Writer {
53 return .{ .context = &self.buffered_writer };
54 }
55};