authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-04 13:49:53+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
loga9c8376305d4c99591432e0ed267fad665bb4f5f
tree91d055c7966263b911f01374b865d03b8cffddc6
parent06ab5a2cd21ecd972477f067bb1d595a9ebef483

std.compress.zstandard: make ZstandardStream decode multiple frames


1 files changed, 97 insertions(+), 53 deletions(-)

lib/std/compress/zstandard.zig+97-53
...@@ -13,6 +13,7 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool...@@ -13,6 +13,7 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
1313
14 allocator: Allocator,14 allocator: Allocator,
15 in_reader: ReaderType,15 in_reader: ReaderType,
16 state: enum { NewFrame, InFrame },
16 decode_state: decompress.block.DecodeState,17 decode_state: decompress.block.DecodeState,
17 frame_context: decompress.FrameContext,18 frame_context: decompress.FrameContext,
18 buffer: RingBuffer,19 buffer: RingBuffer,
...@@ -24,16 +25,43 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool...@@ -24,16 +25,43 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
24 sequence_buffer: []u8,25 sequence_buffer: []u8,
25 checksum: if (verify_checksum) ?u32 else void,26 checksum: if (verify_checksum) ?u32 else void,
2627
27 pub const Error = ReaderType.Error || error{ MalformedBlock, MalformedFrame };28 pub const Error = ReaderType.Error || error{ ChecksumFailure, MalformedBlock, MalformedFrame, OutOfMemory };
2829
29 pub const Reader = std.io.Reader(*Self, Error, read);30 pub const Reader = std.io.Reader(*Self, Error, read);
3031
31 pub fn init(allocator: Allocator, source: ReaderType) !Self {32 pub fn init(allocator: Allocator, source: ReaderType) !Self {
32 switch (try decompress.decodeFrameType(source)) {33 return Self{
33 .skippable => return error.SkippableFrame,34 .allocator = allocator,
35 .in_reader = source,
36 .state = .NewFrame,
37 .decode_state = undefined,
38 .frame_context = undefined,
39 .buffer = undefined,
40 .last_block = undefined,
41 .literal_fse_buffer = undefined,
42 .match_fse_buffer = undefined,
43 .offset_fse_buffer = undefined,
44 .literals_buffer = undefined,
45 .sequence_buffer = undefined,
46 .checksum = undefined,
47 };
48 }
49
50 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, .{});
60 self.state = .NewFrame;
61 },
34 .zstandard => {62 .zstandard => {
35 const frame_context = context: {63 const frame_context = context: {
36 const frame_header = try decompress.decodeZstandardHeader(source);64 const frame_header = try decompress.decodeZstandardHeader(self.in_reader);
37 break :context try decompress.FrameContext.init(65 break :context try decompress.FrameContext.init(
38 frame_header,66 frame_header,
39 window_size_max,67 window_size_max,
...@@ -41,56 +69,58 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool...@@ -41,56 +69,58 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
41 );69 );
42 };70 };
4371
44 const literal_fse_buffer = try allocator.alloc(72 const literal_fse_buffer = try self.allocator.alloc(
45 types.compressed_block.Table.Fse,73 types.compressed_block.Table.Fse,
46 types.compressed_block.table_size_max.literal,74 types.compressed_block.table_size_max.literal,
47 );75 );
48 errdefer allocator.free(literal_fse_buffer);76 errdefer self.allocator.free(literal_fse_buffer);
4977
50 const match_fse_buffer = try allocator.alloc(78 const match_fse_buffer = try self.allocator.alloc(
51 types.compressed_block.Table.Fse,79 types.compressed_block.Table.Fse,
52 types.compressed_block.table_size_max.match,80 types.compressed_block.table_size_max.match,
53 );81 );
54 errdefer allocator.free(match_fse_buffer);82 errdefer self.allocator.free(match_fse_buffer);
5583
56 const offset_fse_buffer = try allocator.alloc(84 const offset_fse_buffer = try self.allocator.alloc(
57 types.compressed_block.Table.Fse,85 types.compressed_block.Table.Fse,
58 types.compressed_block.table_size_max.offset,86 types.compressed_block.table_size_max.offset,
59 );87 );
60 errdefer allocator.free(offset_fse_buffer);88 errdefer self.allocator.free(offset_fse_buffer);
6189
62 const decode_state = decompress.block.DecodeState.init(90 const decode_state = decompress.block.DecodeState.init(
63 literal_fse_buffer,91 literal_fse_buffer,
64 match_fse_buffer,92 match_fse_buffer,
65 offset_fse_buffer,93 offset_fse_buffer,
66 );94 );
67 const buffer = try RingBuffer.init(allocator, frame_context.window_size);95 const buffer = try RingBuffer.init(self.allocator, frame_context.window_size);
6896
69 const literals_data = try allocator.alloc(u8, window_size_max);97 const literals_data = try self.allocator.alloc(u8, window_size_max);
70 errdefer allocator.free(literals_data);98 errdefer self.allocator.free(literals_data);
7199
72 const sequence_data = try allocator.alloc(u8, window_size_max);100 const sequence_data = try self.allocator.alloc(u8, window_size_max);
73 errdefer allocator.free(sequence_data);101 errdefer self.allocator.free(sequence_data);
74102
75 return Self{103 self.literal_fse_buffer = literal_fse_buffer;
76 .allocator = allocator,104 self.match_fse_buffer = match_fse_buffer;
77 .in_reader = source,105 self.offset_fse_buffer = offset_fse_buffer;
78 .decode_state = decode_state,106 self.literals_buffer = literals_data;
79 .frame_context = frame_context,107 self.sequence_buffer = sequence_data;
80 .buffer = buffer,108
81 .checksum = if (verify_checksum) null else {},109 self.buffer = buffer;
82 .last_block = false,110
83 .literal_fse_buffer = literal_fse_buffer,111 self.decode_state = decode_state;
84 .match_fse_buffer = match_fse_buffer,112 self.frame_context = frame_context;
85 .offset_fse_buffer = offset_fse_buffer,113
86 .literals_buffer = literals_data,114 self.checksum = if (verify_checksum) null else {};
87 .sequence_buffer = sequence_data,115 self.last_block = false;
88 };116
117 self.state = .InFrame;
89 },118 },
90 }119 }
91 }120 }
92121
93 pub fn deinit(self: *Self) void {122 pub fn deinit(self: *Self) void {
123 if (self.state == .NewFrame) return;
94 self.allocator.free(self.decode_state.literal_fse_buffer);124 self.allocator.free(self.decode_state.literal_fse_buffer);
95 self.allocator.free(self.decode_state.match_fse_buffer);125 self.allocator.free(self.decode_state.match_fse_buffer);
96 self.allocator.free(self.decode_state.offset_fse_buffer);126 self.allocator.free(self.decode_state.offset_fse_buffer);
...@@ -105,6 +135,19 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool...@@ -105,6 +135,19 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
105135
106 pub fn read(self: *Self, buffer: []u8) Error!usize {136 pub fn read(self: *Self, buffer: []u8) Error!usize {
107 if (buffer.len == 0) return 0;137 if (buffer.len == 0) return 0;
138 while (self.state == .NewFrame) {
139 self.frameInit() catch |err| switch (err) {
140 error.NoBytes => return 0,
141 error.OutOfMemory => return error.OutOfMemory,
142 else => return error.MalformedFrame,
143 };
144 }
145
146 return self.readInner(buffer);
147 }
148
149 fn readInner(self: *Self, buffer: []u8) Error!usize {
150 std.debug.assert(self.state == .InFrame);
108151
109 if (self.buffer.isEmpty() and !self.last_block) {152 if (self.buffer.isEmpty() and !self.last_block) {
110 const header_bytes = self.in_reader.readBytesNoEof(3) catch return error.MalformedFrame;153 const header_bytes = self.in_reader.readBytesNoEof(3) catch return error.MalformedFrame;
...@@ -127,9 +170,15 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool...@@ -127,9 +170,15 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
127 hasher.update(written_slice.first);170 hasher.update(written_slice.first);
128 hasher.update(written_slice.second);171 hasher.update(written_slice.second);
129 }172 }
130 if (block_header.last_block and self.frame_context.has_checksum) {173 if (block_header.last_block) {
131 const checksum = self.in_reader.readIntLittle(u32) catch return error.MalformedFrame;174 if (self.frame_context.has_checksum) {
132 if (verify_checksum) self.checksum = checksum;175 const checksum = self.in_reader.readIntLittle(u32) catch return error.MalformedFrame;
176 if (comptime verify_checksum) {
177 if (self.frame_context.hasher_opt) |*hasher| {
178 if (checksum != decompress.computeChecksum(hasher)) return error.ChecksumFailure;
179 }
180 }
181 }
133 }182 }
134 }183 }
135184
...@@ -138,18 +187,16 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool...@@ -138,18 +187,16 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
138 while (written_count < decoded_data_len and written_count < buffer.len) : (written_count += 1) {187 while (written_count < decoded_data_len and written_count < buffer.len) : (written_count += 1) {
139 buffer[written_count] = self.buffer.read().?;188 buffer[written_count] = self.buffer.read().?;
140 }189 }
141 return written_count;190 if (self.buffer.len() == 0) {
142 }191 self.state = .NewFrame;
143192 self.allocator.free(self.literal_fse_buffer);
144 pub fn verifyChecksum(self: *Self) !bool {193 self.allocator.free(self.match_fse_buffer);
145 if (verify_checksum) {194 self.allocator.free(self.offset_fse_buffer);
146 if (self.checksum) |checksum| {195 self.allocator.free(self.literals_buffer);
147 if (self.frame_context.hasher_opt) |*hasher| {196 self.allocator.free(self.sequence_buffer);
148 return checksum == decompress.computeChecksum(hasher);197 self.buffer.deinit(self.allocator);
149 }
150 }
151 }198 }
152 return true;199 return written_count;
153 }200 }
154 };201 };
155}202}
...@@ -163,7 +210,6 @@ fn testDecompress(data: []const u8) ![]u8 {...@@ -163,7 +210,6 @@ fn testDecompress(data: []const u8) ![]u8 {
163 var stream = try zstandardStream(std.testing.allocator, in_stream.reader());210 var stream = try zstandardStream(std.testing.allocator, in_stream.reader());
164 defer stream.deinit();211 defer stream.deinit();
165 const result = stream.reader().readAllAlloc(std.testing.allocator, std.math.maxInt(usize));212 const result = stream.reader().readAllAlloc(std.testing.allocator, std.math.maxInt(usize));
166 try std.testing.expect(try stream.verifyChecksum());
167 return result;213 return result;
168}214}
169215
...@@ -181,14 +227,12 @@ test "decompression" {...@@ -181,14 +227,12 @@ test "decompression" {
181 var buffer = try std.testing.allocator.alloc(u8, uncompressed.len);227 var buffer = try std.testing.allocator.alloc(u8, uncompressed.len);
182 defer std.testing.allocator.free(buffer);228 defer std.testing.allocator.free(buffer);
183229
184 const res3 = try decompress.decodeFrame(buffer, compressed3, true);230 const res3 = try decompress.decode(buffer, compressed3, true);
185 try std.testing.expectEqual(compressed3.len, res3.read_count);231 try std.testing.expectEqual(uncompressed.len, res3);
186 try std.testing.expectEqual(uncompressed.len, res3.write_count);
187 try std.testing.expectEqualSlices(u8, uncompressed, buffer);232 try std.testing.expectEqualSlices(u8, uncompressed, buffer);
188233
189 const res19 = try decompress.decodeFrame(buffer, compressed19, true);234 const res19 = try decompress.decode(buffer, compressed19, true);
190 try std.testing.expectEqual(compressed19.len, res19.read_count);235 try std.testing.expectEqual(uncompressed.len, res19);
191 try std.testing.expectEqual(uncompressed.len, res19.write_count);
192 try std.testing.expectEqualSlices(u8, uncompressed, buffer);236 try std.testing.expectEqualSlices(u8, uncompressed, buffer);
193237
194 try testReader(compressed3, uncompressed);238 try testReader(compressed3, uncompressed);