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