| ... | @@ -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, |
| 30 | | 31 | |
| 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 | } |
| 56 | | 58 | |
| ... | @@ -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; |
| 114 | | 116 | |
| 115 | self.checksum = if (verify_checksum) null else {}; | 117 | self.checksum = if (verify_checksum) null else {}; |
| | 118 | self.current_frame_decompressed_size = 0; |
| 116 | | 119 | |
| 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 | } |
| 135 | | 138 | |
| 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 | } | | |
| 149 | | 141 | |
| 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 | } |
| 152 | | 159 | |
| 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( |
| 172 | | 179 | |
| 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 | } |
| 195 | | 208 | |
| 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) { |