authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-09 20:15:00+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log31cc4605aba68edc44a31f8383eaa39a906d6ec8
tree0e2b8b02686605858c9373a040b27aeb3df41b75
parent55e6e9409c3c94cca6eb144388104cb03247b045

std.compress.zstandard: fix errors and crashes in ZstandardStream


2 files changed, 30 insertions(+), 30 deletions(-)

lib/std/compress/zstandard.zig+29-30
......@@ -12,12 +12,11 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
1212 const Self = @This();
1313
1414 allocator: Allocator,
15 in_reader: ReaderType,
16 state: enum { NewFrame, InFrame },
15 source: std.io.CountingReader(ReaderType),
16 state: enum { NewFrame, InFrame, LastBlock },
1717 decode_state: decompress.block.DecodeState,
1818 frame_context: decompress.FrameContext,
1919 buffer: RingBuffer,
20 last_block: bool,
2120 literal_fse_buffer: []types.compressed_block.Table.Fse,
2221 match_fse_buffer: []types.compressed_block.Table.Fse,
2322 offset_fse_buffer: []types.compressed_block.Table.Fse,
......@@ -32,12 +31,11 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
3231 pub fn init(allocator: Allocator, source: ReaderType) !Self {
3332 return Self{
3433 .allocator = allocator,
35 .in_reader = source,
34 .source = std.io.countingReader(source),
3635 .state = .NewFrame,
3736 .decode_state = undefined,
3837 .frame_context = undefined,
3938 .buffer = undefined,
40 .last_block = undefined,
4139 .literal_fse_buffer = undefined,
4240 .match_fse_buffer = undefined,
4341 .offset_fse_buffer = undefined,
......@@ -48,22 +46,16 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
4846 }
4947
5048 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, .{});
6053 self.state = .NewFrame;
6154 },
62 .zstandard => {
55 .zstandard => |header| {
6356 const frame_context = context: {
64 const frame_header = try decompress.decodeZstandardHeader(self.in_reader);
6557 break :context try decompress.FrameContext.init(
66 frame_header,
58 header,
6759 window_size_max,
6860 verify_checksum,
6961 );
......@@ -112,7 +104,6 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
112104 self.frame_context = frame_context;
113105
114106 self.checksum = if (verify_checksum) null else {};
115 self.last_block = false;
116107
117108 self.state = .InFrame;
118109 },
......@@ -134,10 +125,14 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
134125 }
135126
136127 pub fn read(self: *Self, buffer: []u8) Error!usize {
128 const initial_count = self.source.bytes_read;
137129 if (buffer.len == 0) return 0;
138130 while (self.state == .NewFrame) {
139131 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,
141136 error.OutOfMemory => return error.OutOfMemory,
142137 else => return error.MalformedFrame,
143138 };
......@@ -147,15 +142,16 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
147142 }
148143
149144 fn readInner(self: *Self, buffer: []u8) Error!usize {
150 std.debug.assert(self.state == .InFrame);
145 std.debug.assert(self.state != .NewFrame);
151146
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;
154150 const block_header = decompress.block.decodeBlockHeader(&header_bytes);
155151
156152 decompress.block.decodeBlockReader(
157153 &self.buffer,
158 self.in_reader,
154 source_reader,
159155 block_header,
160156 &self.decode_state,
161157 self.frame_context.block_size_max,
......@@ -164,15 +160,18 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
164160 ) catch
165161 return error.MalformedBlock;
166162
167 self.last_block = block_header.last_block;
168163 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 }
172170 }
173171 if (block_header.last_block) {
172 self.state = .LastBlock;
174173 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;
176175 if (comptime verify_checksum) {
177176 if (self.frame_context.hasher_opt) |*hasher| {
178177 if (checksum != decompress.computeChecksum(hasher)) return error.ChecksumFailure;
......@@ -187,7 +186,7 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
187186 while (written_count < decoded_data_len and written_count < buffer.len) : (written_count += 1) {
188187 buffer[written_count] = self.buffer.read().?;
189188 }
190 if (self.buffer.len() == 0) {
189 if (self.state == .LastBlock and self.buffer.len() == 0) {
191190 self.state = .NewFrame;
192191 self.allocator.free(self.literal_fse_buffer);
193192 self.allocator.free(self.match_fse_buffer);
......@@ -219,7 +218,7 @@ fn testReader(data: []const u8, comptime expected: []const u8) !void {
219218 try std.testing.expectEqualSlices(u8, expected, buf);
220219}
221220
222test "decompression" {
221test "zstandard decompression" {
223222 const uncompressed = @embedFile("testdata/rfc8478.txt");
224223 const compressed3 = @embedFile("testdata/rfc8478.txt.zst.3");
225224 const compressed19 = @embedFile("testdata/rfc8478.txt.zst.19");
lib/std/compress/zstandard/decode/block.zig+1
......@@ -795,6 +795,7 @@ pub fn decodeBlockReader(
795795 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
796796 switch (block_header.block_type) {
797797 .raw => {
798 if (block_size == 0) return;
798799 const slice = dest.sliceAt(dest.write_index, block_size);
799800 try source.readNoEof(slice.first);
800801 try source.readNoEof(slice.second);