| ... | ... | @@ -12,12 +12,11 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 12 | 12 | const Self = @This(); |
| 13 | 13 | |
| 14 | 14 | allocator: Allocator, |
| 15 | | in_reader: ReaderType, |
| 16 | | state: enum { NewFrame, InFrame }, |
| 15 | source: std.io.CountingReader(ReaderType), |
| 16 | state: enum { NewFrame, InFrame, LastBlock }, |
| 17 | 17 | decode_state: decompress.block.DecodeState, |
| 18 | 18 | frame_context: decompress.FrameContext, |
| 19 | 19 | buffer: RingBuffer, |
| 20 | | last_block: bool, |
| 21 | 20 | literal_fse_buffer: []types.compressed_block.Table.Fse, |
| 22 | 21 | match_fse_buffer: []types.compressed_block.Table.Fse, |
| 23 | 22 | offset_fse_buffer: []types.compressed_block.Table.Fse, |
| ... | ... | @@ -32,12 +31,11 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 32 | 31 | pub fn init(allocator: Allocator, source: ReaderType) !Self { |
| 33 | 32 | return Self{ |
| 34 | 33 | .allocator = allocator, |
| 35 | | .in_reader = source, |
| 34 | .source = std.io.countingReader(source), |
| 36 | 35 | .state = .NewFrame, |
| 37 | 36 | .decode_state = undefined, |
| 38 | 37 | .frame_context = undefined, |
| 39 | 38 | .buffer = undefined, |
| 40 | | .last_block = undefined, |
| 41 | 39 | .literal_fse_buffer = undefined, |
| 42 | 40 | .match_fse_buffer = undefined, |
| 43 | 41 | .offset_fse_buffer = undefined, |
| ... | ... | @@ -48,22 +46,16 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 48 | 46 | } |
| 49 | 47 | |
| 50 | 48 | fn frameInit(self: *Self) !void { |
| 51 | | var bytes: [4]u8 = undefined; |
| 52 | | const bytes_read = try self.in_reader.readAll(&bytes); |
| 53 | | if (bytes_read == 0) return error.NoBytes; |
| 54 | | if (bytes_read < 4) return error.EndOfStream; |
| 55 | | const frame_type = try decompress.frameType(std.mem.readIntLittle(u32, &bytes)); |
| 56 | | switch (frame_type) { |
| 57 | | .skippable => { |
| 58 | | const size = try self.in_reader.readIntLittle(u32); |
| 59 | | try self.in_reader.skipBytes(size, .{}); |
| 49 | const source_reader = self.source.reader(); |
| 50 | switch (try decompress.decodeFrameHeader(source_reader)) { |
| 51 | .skippable => |header| { |
| 52 | try source_reader.skipBytes(header.frame_size, .{}); |
| 60 | 53 | self.state = .NewFrame; |
| 61 | 54 | }, |
| 62 | | .zstandard => { |
| 55 | .zstandard => |header| { |
| 63 | 56 | const frame_context = context: { |
| 64 | | const frame_header = try decompress.decodeZstandardHeader(self.in_reader); |
| 65 | 57 | break :context try decompress.FrameContext.init( |
| 66 | | frame_header, |
| 58 | header, |
| 67 | 59 | window_size_max, |
| 68 | 60 | verify_checksum, |
| 69 | 61 | ); |
| ... | ... | @@ -112,7 +104,6 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 112 | 104 | self.frame_context = frame_context; |
| 113 | 105 | |
| 114 | 106 | self.checksum = if (verify_checksum) null else {}; |
| 115 | | self.last_block = false; |
| 116 | 107 | |
| 117 | 108 | self.state = .InFrame; |
| 118 | 109 | }, |
| ... | ... | @@ -134,10 +125,14 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 134 | 125 | } |
| 135 | 126 | |
| 136 | 127 | pub fn read(self: *Self, buffer: []u8) Error!usize { |
| 128 | const initial_count = self.source.bytes_read; |
| 137 | 129 | if (buffer.len == 0) return 0; |
| 138 | 130 | while (self.state == .NewFrame) { |
| 139 | 131 | self.frameInit() catch |err| switch (err) { |
| 140 | | error.NoBytes => return 0, |
| 132 | error.EndOfStream => return if (self.source.bytes_read == initial_count) |
| 133 | 0 |
| 134 | else |
| 135 | error.MalformedFrame, |
| 141 | 136 | error.OutOfMemory => return error.OutOfMemory, |
| 142 | 137 | else => return error.MalformedFrame, |
| 143 | 138 | }; |
| ... | ... | @@ -147,15 +142,16 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 147 | 142 | } |
| 148 | 143 | |
| 149 | 144 | fn readInner(self: *Self, buffer: []u8) Error!usize { |
| 150 | | std.debug.assert(self.state == .InFrame); |
| 145 | std.debug.assert(self.state != .NewFrame); |
| 151 | 146 | |
| 152 | | if (self.buffer.isEmpty() and !self.last_block) { |
| 153 | | const header_bytes = self.in_reader.readBytesNoEof(3) catch return error.MalformedFrame; |
| 147 | const source_reader = self.source.reader(); |
| 148 | while (self.buffer.isEmpty() and self.state != .LastBlock) { |
| 149 | const header_bytes = source_reader.readBytesNoEof(3) catch return error.MalformedFrame; |
| 154 | 150 | const block_header = decompress.block.decodeBlockHeader(&header_bytes); |
| 155 | 151 | |
| 156 | 152 | decompress.block.decodeBlockReader( |
| 157 | 153 | &self.buffer, |
| 158 | | self.in_reader, |
| 154 | source_reader, |
| 159 | 155 | block_header, |
| 160 | 156 | &self.decode_state, |
| 161 | 157 | self.frame_context.block_size_max, |
| ... | ... | @@ -164,15 +160,18 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 164 | 160 | ) catch |
| 165 | 161 | return error.MalformedBlock; |
| 166 | 162 | |
| 167 | | self.last_block = block_header.last_block; |
| 168 | 163 | if (self.frame_context.hasher_opt) |*hasher| { |
| 169 | | const written_slice = self.buffer.sliceLast(self.buffer.len()); |
| 170 | | hasher.update(written_slice.first); |
| 171 | | hasher.update(written_slice.second); |
| 164 | const size = self.buffer.len(); |
| 165 | if (size > 0) { |
| 166 | const written_slice = self.buffer.sliceLast(size); |
| 167 | hasher.update(written_slice.first); |
| 168 | hasher.update(written_slice.second); |
| 169 | } |
| 172 | 170 | } |
| 173 | 171 | if (block_header.last_block) { |
| 172 | self.state = .LastBlock; |
| 174 | 173 | if (self.frame_context.has_checksum) { |
| 175 | | const checksum = self.in_reader.readIntLittle(u32) catch return error.MalformedFrame; |
| 174 | const checksum = source_reader.readIntLittle(u32) catch return error.MalformedFrame; |
| 176 | 175 | if (comptime verify_checksum) { |
| 177 | 176 | if (self.frame_context.hasher_opt) |*hasher| { |
| 178 | 177 | if (checksum != decompress.computeChecksum(hasher)) return error.ChecksumFailure; |
| ... | ... | @@ -187,7 +186,7 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool |
| 187 | 186 | while (written_count < decoded_data_len and written_count < buffer.len) : (written_count += 1) { |
| 188 | 187 | buffer[written_count] = self.buffer.read().?; |
| 189 | 188 | } |
| 190 | | if (self.buffer.len() == 0) { |
| 189 | if (self.state == .LastBlock and self.buffer.len() == 0) { |
| 191 | 190 | self.state = .NewFrame; |
| 192 | 191 | self.allocator.free(self.literal_fse_buffer); |
| 193 | 192 | self.allocator.free(self.match_fse_buffer); |
| ... | ... | @@ -219,7 +218,7 @@ fn testReader(data: []const u8, comptime expected: []const u8) !void { |
| 219 | 218 | try std.testing.expectEqualSlices(u8, expected, buf); |
| 220 | 219 | } |
| 221 | 220 | |
| 222 | | test "decompression" { |
| 221 | test "zstandard decompression" { |
| 223 | 222 | const uncompressed = @embedFile("testdata/rfc8478.txt"); |
| 224 | 223 | const compressed3 = @embedFile("testdata/rfc8478.txt.zst.3"); |
| 225 | 224 | const compressed19 = @embedFile("testdata/rfc8478.txt.zst.19"); |