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...@@ -12,12 +12,11 @@ pub fn ZstandardStream(comptime ReaderType: type, comptime verify_checksum: bool
12 const Self = @This();12 const Self = @This();
1313
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 }
4947
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;
113105
114 self.checksum = if (verify_checksum) null else {};106 self.checksum = if (verify_checksum) null else {};
115 self.last_block = false;
116107
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 }
135126
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 }
148143
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);
151146
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);
155151
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 ) catch160 ) catch
165 return error.MalformedBlock;161 return error.MalformedBlock;
166162
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}
221220
222test "decompression" {221test "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");
lib/std/compress/zstandard/decode/block.zig+1
...@@ -795,6 +795,7 @@ pub fn decodeBlockReader(...@@ -795,6 +795,7 @@ pub fn decodeBlockReader(
795 if (block_size_max < block_size) return error.BlockSizeOverMaximum;795 if (block_size_max < block_size) return error.BlockSizeOverMaximum;
796 switch (block_header.block_type) {796 switch (block_header.block_type) {
797 .raw => {797 .raw => {
798 if (block_size == 0) return;
798 const slice = dest.sliceAt(dest.write_index, block_size);799 const slice = dest.sliceAt(dest.write_index, block_size);
799 try source.readNoEof(slice.first);800 try source.readNoEof(slice.first);
800 try source.readNoEof(slice.second);801 try source.readNoEof(slice.second);