authorgravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-12 13:05:34+11:00
committergravatar for 4678790+dweiller@users.noreply.github.comDominic <4678790+dweiller@users.noreply.github.com> 2023-02-20 09:09:06+11:00
log476d2fe1fa917a3b7a92cf155c779eb975c906e6
tree7774215bbe52863c593b2457b543c1a8a5a17ca5
parent373d8ef26edca9d16111ae41f960a44ead6ea2c8

std.compress.zstandard: fix zstandardStream finishing early


1 files changed, 27 insertions(+), 14 deletions(-)

lib/std/compress/zstandard.zig+27-14
...@@ -27,6 +27,7 @@ pub fn ZstandardStream(...@@ -27,6 +27,7 @@ pub fn ZstandardStream(
27 literals_buffer: []u8,27 literals_buffer: []u8,
28 sequence_buffer: []u8,28 sequence_buffer: []u8,
29 checksum: if (verify_checksum) ?u32 else void,29 checksum: if (verify_checksum) ?u32 else void,
30 current_frame_decompressed_size: usize,
3031
31 pub const Error = ReaderType.Error || error{32 pub const Error = ReaderType.Error || error{
32 ChecksumFailure,33 ChecksumFailure,
...@@ -51,6 +52,7 @@ pub fn ZstandardStream(...@@ -51,6 +52,7 @@ pub fn ZstandardStream(
51 .literals_buffer = undefined,52 .literals_buffer = undefined,
52 .sequence_buffer = undefined,53 .sequence_buffer = undefined,
53 .checksum = undefined,54 .checksum = undefined,
55 .current_frame_decompressed_size = undefined,
54 };56 };
55 }57 }
5658
...@@ -113,6 +115,7 @@ pub fn ZstandardStream(...@@ -113,6 +115,7 @@ pub fn ZstandardStream(
113 self.frame_context = frame_context;115 self.frame_context = frame_context;
114116
115 self.checksum = if (verify_checksum) null else {};117 self.checksum = if (verify_checksum) null else {};
118 self.current_frame_decompressed_size = 0;
116119
117 self.state = .InFrame;120 self.state = .InFrame;
118 },121 },
...@@ -134,20 +137,24 @@ pub fn ZstandardStream(...@@ -134,20 +137,24 @@ pub fn ZstandardStream(
134 }137 }
135138
136 pub fn read(self: *Self, buffer: []u8) Error!usize {139 pub fn read(self: *Self, buffer: []u8) Error!usize {
137 const initial_count = self.source.bytes_read;
138 if (buffer.len == 0) return 0;140 if (buffer.len == 0) return 0;
139 while (self.state == .NewFrame) {
140 self.frameInit() catch |err| switch (err) {
141 error.EndOfStream => return if (self.source.bytes_read == initial_count)
142 0
143 else
144 error.MalformedFrame,
145 error.OutOfMemory => return error.OutOfMemory,
146 else => return error.MalformedFrame,
147 };
148 }
149141
150 return self.readInner(buffer);142 var size: usize = 0;
143 while (size == 0) {
144 while (self.state == .NewFrame) {
145 const initial_count = self.source.bytes_read;
146 self.frameInit() catch |err| switch (err) {
147 error.EndOfStream => return if (self.source.bytes_read == initial_count)
148 0
149 else
150 error.MalformedFrame,
151 error.OutOfMemory => return error.OutOfMemory,
152 else => return error.MalformedFrame,
153 };
154 }
155 size = try self.readInner(buffer);
156 }
157 return size;
151 }158 }
152159
153 fn readInner(self: *Self, buffer: []u8) Error!usize {160 fn readInner(self: *Self, buffer: []u8) Error!usize {
...@@ -172,6 +179,7 @@ pub fn ZstandardStream(...@@ -172,6 +179,7 @@ pub fn ZstandardStream(
172179
173 if (self.frame_context.hasher_opt) |*hasher| {180 if (self.frame_context.hasher_opt) |*hasher| {
174 const size = self.buffer.len();181 const size = self.buffer.len();
182 self.current_frame_decompressed_size += size;
175 if (size > 0) {183 if (size > 0) {
176 const written_slice = self.buffer.sliceLast(size);184 const written_slice = self.buffer.sliceLast(size);
177 hasher.update(written_slice.first);185 hasher.update(written_slice.first);
...@@ -190,12 +198,17 @@ pub fn ZstandardStream(...@@ -190,12 +198,17 @@ pub fn ZstandardStream(
190 }198 }
191 }199 }
192 }200 }
201 if (self.frame_context.content_size) |content_size| {
202 if (content_size != self.current_frame_decompressed_size) {
203 return error.MalformedFrame;
204 }
205 }
193 }206 }
194 }207 }
195208
196 const decoded_data_len = self.buffer.len();209 const size = @min(self.buffer.len(), buffer.len);
197 var count: usize = 0;210 var count: usize = 0;
198 while (count < decoded_data_len and count < buffer.len) : (count += 1) {211 while (count < size) : (count += 1) {
199 buffer[count] = self.buffer.read().?;212 buffer[count] = self.buffer.read().?;
200 }213 }
201 if (self.state == .LastBlock and self.buffer.len() == 0) {214 if (self.state == .LastBlock and self.buffer.len() == 0) {