| ... | @@ -22,7 +22,7 @@ fn isSkippableMagic(magic: u32) bool { | ... | @@ -22,7 +22,7 @@ fn isSkippableMagic(magic: u32) bool { |
| 22 | /// if the the frame is skippable, `null` for Zstanndard frames that do not | 22 | /// if the the frame is skippable, `null` for Zstanndard frames that do not |
| 23 | /// declare their content size. Returns `UnusedBitSet` and `ReservedBitSet` | 23 | /// declare their content size. Returns `UnusedBitSet` and `ReservedBitSet` |
| 24 | /// errors if the respective bits of the the frame descriptor are set. | 24 | /// errors if the respective bits of the the frame descriptor are set. |
| 25 | pub fn getFrameDecompressedSize(src: []const u8) !?u64 { | 25 | pub fn getFrameDecompressedSize(src: []const u8) (InvalidBit || error{BadMagic})!?u64 { |
| 26 | switch (try frameType(src)) { | 26 | switch (try frameType(src)) { |
| 27 | .zstandard => { | 27 | .zstandard => { |
| 28 | const header = try decodeZStandardHeader(src[4..], null); | 28 | const header = try decodeZStandardHeader(src[4..], null); |
| ... | @@ -52,7 +52,11 @@ const ReadWriteCount = struct { | ... | @@ -52,7 +52,11 @@ const ReadWriteCount = struct { |
| 52 | | 52 | |
| 53 | /// Decodes the frame at the start of `src` into `dest`. Returns the number of | 53 | /// Decodes the frame at the start of `src` into `dest`. Returns the number of |
| 54 | /// bytes read from `src` and written to `dest`. | 54 | /// bytes read from `src` and written to `dest`. |
| 55 | pub fn decodeFrame(dest: []u8, src: []const u8, verify_checksum: bool) !ReadWriteCount { | 55 | pub fn decodeFrame( |
| | 56 | dest: []u8, |
| | 57 | src: []const u8, |
| | 58 | verify_checksum: bool, |
| | 59 | ) (error{ UnknownContentSizeUnsupported, ContentTooLarge, BadMagic } || FrameError)!ReadWriteCount { |
| 56 | return switch (try frameType(src)) { | 60 | return switch (try frameType(src)) { |
| 57 | .zstandard => decodeZStandardFrame(dest, src, verify_checksum), | 61 | .zstandard => decodeZStandardFrame(dest, src, verify_checksum), |
| 58 | .skippable => ReadWriteCount{ | 62 | .skippable => ReadWriteCount{ |
| ... | @@ -100,7 +104,7 @@ pub const DecodeState = struct { | ... | @@ -100,7 +104,7 @@ pub const DecodeState = struct { |
| 100 | src: []const u8, | 104 | src: []const u8, |
| 101 | literals: LiteralsSection, | 105 | literals: LiteralsSection, |
| 102 | sequences_header: SequencesSection.Header, | 106 | sequences_header: SequencesSection.Header, |
| 103 | ) !usize { | 107 | ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize { |
| 104 | if (literals.huffman_tree) |tree| { | 108 | if (literals.huffman_tree) |tree| { |
| 105 | self.huffman_tree = tree; | 109 | self.huffman_tree = tree; |
| 106 | } else if (literals.header.block_type == .treeless and self.huffman_tree == null) { | 110 | } else if (literals.header.block_type == .treeless and self.huffman_tree == null) { |
| ... | @@ -145,7 +149,7 @@ pub const DecodeState = struct { | ... | @@ -145,7 +149,7 @@ pub const DecodeState = struct { |
| 145 | | 149 | |
| 146 | /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream` | 150 | /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream` |
| 147 | /// if `bit_reader` does not contain enough bits. | 151 | /// if `bit_reader` does not contain enough bits. |
| 148 | pub fn readInitialFseState(self: *DecodeState, bit_reader: anytype) !void { | 152 | pub fn readInitialFseState(self: *DecodeState, bit_reader: *ReverseBitReader) error{EndOfStream}!void { |
| 149 | self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log); | 153 | self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log); |
| 150 | self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log); | 154 | self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log); |
| 151 | self.match.state = try bit_reader.readBitsNoEof(u9, self.match.accuracy_log); | 155 | self.match.state = try bit_reader.readBitsNoEof(u9, self.match.accuracy_log); |
| ... | @@ -169,7 +173,11 @@ pub const DecodeState = struct { | ... | @@ -169,7 +173,11 @@ pub const DecodeState = struct { |
| 169 | | 173 | |
| 170 | const DataType = enum { offset, match, literal }; | 174 | const DataType = enum { offset, match, literal }; |
| 171 | | 175 | |
| 172 | fn updateState(self: *DecodeState, comptime choice: DataType, bit_reader: anytype) !void { | 176 | fn updateState( |
| | 177 | self: *DecodeState, |
| | 178 | comptime choice: DataType, |
| | 179 | bit_reader: *ReverseBitReader, |
| | 180 | ) error{ MalformedFseBits, EndOfStream }!void { |
| 173 | switch (@field(self, @tagName(choice)).table) { | 181 | switch (@field(self, @tagName(choice)).table) { |
| 174 | .rle => {}, | 182 | .rle => {}, |
| 175 | .fse => |table| { | 183 | .fse => |table| { |
| ... | @@ -185,17 +193,27 @@ pub const DecodeState = struct { | ... | @@ -185,17 +193,27 @@ pub const DecodeState = struct { |
| 185 | } | 193 | } |
| 186 | } | 194 | } |
| 187 | | 195 | |
| | 196 | const FseTableError = error{ |
| | 197 | MalformedFseTable, |
| | 198 | MalformedAccuracyLog, |
| | 199 | RepeatModeFirst, |
| | 200 | EndOfStream, |
| | 201 | }; |
| | 202 | |
| 188 | fn updateFseTable( | 203 | fn updateFseTable( |
| 189 | self: *DecodeState, | 204 | self: *DecodeState, |
| 190 | src: []const u8, | 205 | src: []const u8, |
| 191 | comptime choice: DataType, | 206 | comptime choice: DataType, |
| 192 | mode: SequencesSection.Header.Mode, | 207 | mode: SequencesSection.Header.Mode, |
| 193 | ) !usize { | 208 | ) FseTableError!usize { |
| 194 | const field_name = @tagName(choice); | 209 | const field_name = @tagName(choice); |
| 195 | switch (mode) { | 210 | switch (mode) { |
| 196 | .predefined => { | 211 | .predefined => { |
| 197 | @field(self, field_name).accuracy_log = @field(types.compressed_block.default_accuracy_log, field_name); | 212 | @field(self, field_name).accuracy_log = |
| 198 | @field(self, field_name).table = @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table"); | 213 | @field(types.compressed_block.default_accuracy_log, field_name); |
| | 214 | |
| | 215 | @field(self, field_name).table = |
| | 216 | @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table"); |
| 199 | return 0; | 217 | return 0; |
| 200 | }, | 218 | }, |
| 201 | .rle => { | 219 | .rle => { |
| ... | @@ -214,9 +232,11 @@ pub const DecodeState = struct { | ... | @@ -214,9 +232,11 @@ pub const DecodeState = struct { |
| 214 | @field(types.compressed_block.table_accuracy_log_max, field_name), | 232 | @field(types.compressed_block.table_accuracy_log_max, field_name), |
| 215 | @field(self, field_name ++ "_fse_buffer"), | 233 | @field(self, field_name ++ "_fse_buffer"), |
| 216 | ); | 234 | ); |
| 217 | @field(self, field_name).table = .{ .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size] }; | 235 | @field(self, field_name).table = .{ |
| | 236 | .fse = @field(self, field_name ++ "_fse_buffer")[0..table_size], |
| | 237 | }; |
| 218 | @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size); | 238 | @field(self, field_name).accuracy_log = std.math.log2_int_ceil(usize, table_size); |
| 219 | return std.math.cast(usize, counting_reader.bytes_read) orelse return error.MalformedFseTable; | 239 | return std.math.cast(usize, counting_reader.bytes_read) orelse error.MalformedFseTable; |
| 220 | }, | 240 | }, |
| 221 | .repeat => return if (self.fse_tables_undefined) error.RepeatModeFirst else 0, | 241 | .repeat => return if (self.fse_tables_undefined) error.RepeatModeFirst else 0, |
| 222 | } | 242 | } |
| ... | @@ -228,7 +248,10 @@ pub const DecodeState = struct { | ... | @@ -228,7 +248,10 @@ pub const DecodeState = struct { |
| 228 | offset: u32, | 248 | offset: u32, |
| 229 | }; | 249 | }; |
| 230 | | 250 | |
| 231 | fn nextSequence(self: *DecodeState, bit_reader: anytype) !Sequence { | 251 | fn nextSequence( |
| | 252 | self: *DecodeState, |
| | 253 | bit_reader: *ReverseBitReader, |
| | 254 | ) error{ OffsetCodeTooLarge, EndOfStream }!Sequence { |
| 232 | const raw_code = self.getCode(.offset); | 255 | const raw_code = self.getCode(.offset); |
| 233 | const offset_code = std.math.cast(u5, raw_code) orelse { | 256 | const offset_code = std.math.cast(u5, raw_code) orelse { |
| 234 | return error.OffsetCodeTooLarge; | 257 | return error.OffsetCodeTooLarge; |
| ... | @@ -272,7 +295,7 @@ pub const DecodeState = struct { | ... | @@ -272,7 +295,7 @@ pub const DecodeState = struct { |
| 272 | write_pos: usize, | 295 | write_pos: usize, |
| 273 | literals: LiteralsSection, | 296 | literals: LiteralsSection, |
| 274 | sequence: Sequence, | 297 | sequence: Sequence, |
| 275 | ) !void { | 298 | ) (error{MalformedSequence} || DecodeLiteralsError)!void { |
| 276 | if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence; | 299 | if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence; |
| 277 | | 300 | |
| 278 | try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length); | 301 | try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length); |
| ... | @@ -288,16 +311,23 @@ pub const DecodeState = struct { | ... | @@ -288,16 +311,23 @@ pub const DecodeState = struct { |
| 288 | dest: *RingBuffer, | 311 | dest: *RingBuffer, |
| 289 | literals: LiteralsSection, | 312 | literals: LiteralsSection, |
| 290 | sequence: Sequence, | 313 | sequence: Sequence, |
| 291 | ) !void { | 314 | ) (error{MalformedSequence} || DecodeLiteralsError)!void { |
| 292 | if (sequence.offset > dest.data.len) return error.MalformedSequence; | 315 | if (sequence.offset > dest.data.len) return error.MalformedSequence; |
| 293 | | 316 | |
| 294 | try self.decodeLiteralsRingBuffer(dest, literals, sequence.literal_length); | 317 | try self.decodeLiteralsRingBuffer(dest, literals, sequence.literal_length); |
| 295 | const copy_slice = dest.sliceAt(dest.write_index + dest.data.len - sequence.offset, sequence.match_length); | 318 | const copy_start = dest.write_index + dest.data.len - sequence.offset; |
| | 319 | const copy_slice = dest.sliceAt(copy_start, sequence.match_length); |
| 296 | // TODO: would std.mem.copy and figuring out dest slice be better/faster? | 320 | // TODO: would std.mem.copy and figuring out dest slice be better/faster? |
| 297 | for (copy_slice.first) |b| dest.writeAssumeCapacity(b); | 321 | for (copy_slice.first) |b| dest.writeAssumeCapacity(b); |
| 298 | for (copy_slice.second) |b| dest.writeAssumeCapacity(b); | 322 | for (copy_slice.second) |b| dest.writeAssumeCapacity(b); |
| 299 | } | 323 | } |
| 300 | | 324 | |
| | 325 | const DecodeSequenceError = error{ |
| | 326 | OffsetCodeTooLarge, |
| | 327 | EndOfStream, |
| | 328 | MalformedSequence, |
| | 329 | MalformedFseBits, |
| | 330 | } || DecodeLiteralsError; |
| 301 | /// Decode one sequence from `bit_reader` into `dest`, written starting at | 331 | /// Decode one sequence from `bit_reader` into `dest`, written starting at |
| 302 | /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns | 332 | /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns |
| 303 | /// `error.MalformedSequence` error if the decompressed sequence would be longer | 333 | /// `error.MalformedSequence` error if the decompressed sequence would be longer |
| ... | @@ -311,10 +341,10 @@ pub const DecodeState = struct { | ... | @@ -311,10 +341,10 @@ pub const DecodeState = struct { |
| 311 | dest: []u8, | 341 | dest: []u8, |
| 312 | write_pos: usize, | 342 | write_pos: usize, |
| 313 | literals: LiteralsSection, | 343 | literals: LiteralsSection, |
| 314 | bit_reader: anytype, | 344 | bit_reader: *ReverseBitReader, |
| 315 | sequence_size_limit: usize, | 345 | sequence_size_limit: usize, |
| 316 | last_sequence: bool, | 346 | last_sequence: bool, |
| 317 | ) !usize { | 347 | ) DecodeSequenceError!usize { |
| 318 | const sequence = try self.nextSequence(bit_reader); | 348 | const sequence = try self.nextSequence(bit_reader); |
| 319 | const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length; | 349 | const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length; |
| 320 | if (sequence_length > sequence_size_limit) return error.MalformedSequence; | 350 | if (sequence_length > sequence_size_limit) return error.MalformedSequence; |
| ... | @@ -336,7 +366,7 @@ pub const DecodeState = struct { | ... | @@ -336,7 +366,7 @@ pub const DecodeState = struct { |
| 336 | bit_reader: anytype, | 366 | bit_reader: anytype, |
| 337 | sequence_size_limit: usize, | 367 | sequence_size_limit: usize, |
| 338 | last_sequence: bool, | 368 | last_sequence: bool, |
| 339 | ) !usize { | 369 | ) DecodeSequenceError!usize { |
| 340 | const sequence = try self.nextSequence(bit_reader); | 370 | const sequence = try self.nextSequence(bit_reader); |
| 341 | const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length; | 371 | const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length; |
| 342 | if (sequence_length > sequence_size_limit) return error.MalformedSequence; | 372 | if (sequence_length > sequence_size_limit) return error.MalformedSequence; |
| ... | @@ -350,26 +380,63 @@ pub const DecodeState = struct { | ... | @@ -350,26 +380,63 @@ pub const DecodeState = struct { |
| 350 | return sequence_length; | 380 | return sequence_length; |
| 351 | } | 381 | } |
| 352 | | 382 | |
| 353 | fn nextLiteralMultiStream(self: *DecodeState, literals: LiteralsSection) !void { | 383 | fn nextLiteralMultiStream( |
| | 384 | self: *DecodeState, |
| | 385 | literals: LiteralsSection, |
| | 386 | ) error{BitStreamHasNoStartBit}!void { |
| 354 | self.literal_stream_index += 1; | 387 | self.literal_stream_index += 1; |
| 355 | try self.initLiteralStream(literals.streams.four[self.literal_stream_index]); | 388 | try self.initLiteralStream(literals.streams.four[self.literal_stream_index]); |
| 356 | } | 389 | } |
| 357 | | 390 | |
| 358 | fn initLiteralStream(self: *DecodeState, bytes: []const u8) !void { | 391 | fn initLiteralStream(self: *DecodeState, bytes: []const u8) error{BitStreamHasNoStartBit}!void { |
| 359 | try self.literal_stream_reader.init(bytes); | 392 | try self.literal_stream_reader.init(bytes); |
| 360 | } | 393 | } |
| 361 | | 394 | |
| | 395 | const LiteralBitsError = error{ |
| | 396 | BitStreamHasNoStartBit, |
| | 397 | UnexpectedEndOfLiteralStream, |
| | 398 | }; |
| | 399 | fn readLiteralsBits( |
| | 400 | self: *DecodeState, |
| | 401 | comptime T: type, |
| | 402 | bit_count_to_read: usize, |
| | 403 | literals: LiteralsSection, |
| | 404 | ) LiteralBitsError!T { |
| | 405 | return self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch bits: { |
| | 406 | if (literals.streams == .four and self.literal_stream_index < 3) { |
| | 407 | try self.nextLiteralMultiStream(literals); |
| | 408 | break :bits self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |
| | 409 | return error.UnexpectedEndOfLiteralStream; |
| | 410 | } else { |
| | 411 | return error.UnexpectedEndOfLiteralStream; |
| | 412 | } |
| | 413 | }; |
| | 414 | } |
| | 415 | |
| | 416 | const DecodeLiteralsError = error{ |
| | 417 | MalformedLiteralsLength, |
| | 418 | PrefixNotFound, |
| | 419 | } || LiteralBitsError; |
| | 420 | |
| 362 | /// Decode `len` bytes of literals into `dest`. `literals` should be the | 421 | /// Decode `len` bytes of literals into `dest`. `literals` should be the |
| 363 | /// `LiteralsSection` that was passed to `prepare()`. Returns | 422 | /// `LiteralsSection` that was passed to `prepare()`. Returns |
| 364 | /// `error.MalformedLiteralsLength` if the number of literal bytes decoded by | 423 | /// `error.MalformedLiteralsLength` if the number of literal bytes decoded by |
| 365 | /// `self` plus `len` is greater than the regenerated size of `literals`. | 424 | /// `self` plus `len` is greater than the regenerated size of `literals`. |
| 366 | /// Returns `error.UnexpectedEndOfLiteralStream` and `error.PrefixNotFound` if | 425 | /// Returns `error.UnexpectedEndOfLiteralStream` and `error.PrefixNotFound` if |
| 367 | /// there are problems decoding Huffman compressed literals. | 426 | /// there are problems decoding Huffman compressed literals. |
| 368 | pub fn decodeLiteralsSlice(self: *DecodeState, dest: []u8, literals: LiteralsSection, len: usize) !void { | 427 | pub fn decodeLiteralsSlice( |
| 369 | if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength; | 428 | self: *DecodeState, |
| | 429 | dest: []u8, |
| | 430 | literals: LiteralsSection, |
| | 431 | len: usize, |
| | 432 | ) DecodeLiteralsError!void { |
| | 433 | if (self.literal_written_count + len > literals.header.regenerated_size) |
| | 434 | return error.MalformedLiteralsLength; |
| | 435 | |
| 370 | switch (literals.header.block_type) { | 436 | switch (literals.header.block_type) { |
| 371 | .raw => { | 437 | .raw => { |
| 372 | const literal_data = literals.streams.one[self.literal_written_count .. self.literal_written_count + len]; | 438 | const literals_end = self.literal_written_count + len; |
| | 439 | const literal_data = literals.streams.one[self.literal_written_count..literals_end]; |
| 373 | std.mem.copy(u8, dest, literal_data); | 440 | std.mem.copy(u8, dest, literal_data); |
| 374 | self.literal_written_count += len; | 441 | self.literal_written_count += len; |
| 375 | }, | 442 | }, |
| ... | @@ -395,15 +462,7 @@ pub const DecodeState = struct { | ... | @@ -395,15 +462,7 @@ pub const DecodeState = struct { |
| 395 | while (i < len) : (i += 1) { | 462 | while (i < len) : (i += 1) { |
| 396 | var prefix: u16 = 0; | 463 | var prefix: u16 = 0; |
| 397 | while (true) { | 464 | while (true) { |
| 398 | const new_bits = self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |err| | 465 | const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals); |
| 399 | switch (err) { | | |
| 400 | error.EndOfStream => if (literals.streams == .four and self.literal_stream_index < 3) bits: { | | |
| 401 | try self.nextLiteralMultiStream(literals); | | |
| 402 | break :bits try self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read); | | |
| 403 | } else { | | |
| 404 | return error.UnexpectedEndOfLiteralStream; | | |
| 405 | }, | | |
| 406 | }; | | |
| 407 | prefix <<= bit_count_to_read; | 466 | prefix <<= bit_count_to_read; |
| 408 | prefix |= new_bits; | 467 | prefix |= new_bits; |
| 409 | bits_read += bit_count_to_read; | 468 | bits_read += bit_count_to_read; |
| ... | @@ -434,11 +493,19 @@ pub const DecodeState = struct { | ... | @@ -434,11 +493,19 @@ pub const DecodeState = struct { |
| 434 | } | 493 | } |
| 435 | | 494 | |
| 436 | /// Decode literals into `dest`; see `decodeLiteralsSlice()`. | 495 | /// Decode literals into `dest`; see `decodeLiteralsSlice()`. |
| 437 | pub fn decodeLiteralsRingBuffer(self: *DecodeState, dest: *RingBuffer, literals: LiteralsSection, len: usize) !void { | 496 | pub fn decodeLiteralsRingBuffer( |
| 438 | if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength; | 497 | self: *DecodeState, |
| | 498 | dest: *RingBuffer, |
| | 499 | literals: LiteralsSection, |
| | 500 | len: usize, |
| | 501 | ) DecodeLiteralsError!void { |
| | 502 | if (self.literal_written_count + len > literals.header.regenerated_size) |
| | 503 | return error.MalformedLiteralsLength; |
| | 504 | |
| 439 | switch (literals.header.block_type) { | 505 | switch (literals.header.block_type) { |
| 440 | .raw => { | 506 | .raw => { |
| 441 | const literal_data = literals.streams.one[self.literal_written_count .. self.literal_written_count + len]; | 507 | const literals_end = self.literal_written_count + len; |
| | 508 | const literal_data = literals.streams.one[self.literal_written_count..literals_end]; |
| 442 | dest.writeSliceAssumeCapacity(literal_data); | 509 | dest.writeSliceAssumeCapacity(literal_data); |
| 443 | self.literal_written_count += len; | 510 | self.literal_written_count += len; |
| 444 | }, | 511 | }, |
| ... | @@ -464,15 +531,7 @@ pub const DecodeState = struct { | ... | @@ -464,15 +531,7 @@ pub const DecodeState = struct { |
| 464 | while (i < len) : (i += 1) { | 531 | while (i < len) : (i += 1) { |
| 465 | var prefix: u16 = 0; | 532 | var prefix: u16 = 0; |
| 466 | while (true) { | 533 | while (true) { |
| 467 | const new_bits = self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |err| | 534 | const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals); |
| 468 | switch (err) { | | |
| 469 | error.EndOfStream => if (literals.streams == .four and self.literal_stream_index < 3) bits: { | | |
| 470 | try self.nextLiteralMultiStream(literals); | | |
| 471 | break :bits try self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read); | | |
| 472 | } else { | | |
| 473 | return error.UnexpectedEndOfLiteralStream; | | |
| 474 | }, | | |
| 475 | }; | | |
| 476 | prefix <<= bit_count_to_read; | 535 | prefix <<= bit_count_to_read; |
| 477 | prefix |= new_bits; | 536 | prefix |= new_bits; |
| 478 | bits_read += bit_count_to_read; | 537 | bits_read += bit_count_to_read; |
| ... | @@ -514,6 +573,11 @@ const literal_table_size_max = 1 << types.compressed_block.table_accuracy_log_ma | ... | @@ -514,6 +573,11 @@ const literal_table_size_max = 1 << types.compressed_block.table_accuracy_log_ma |
| 514 | const match_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match; | 573 | const match_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match; |
| 515 | const offset_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match; | 574 | const offset_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match; |
| 516 | | 575 | |
| | 576 | const FrameError = error{ |
| | 577 | DictionaryIdFlagUnsupported, |
| | 578 | ChecksumFailure, |
| | 579 | } || InvalidBit || DecodeBlockError; |
| | 580 | |
| 517 | /// Decode a Zstandard frame from `src` into `dest`, returning the number of | 581 | /// Decode a Zstandard frame from `src` into `dest`, returning the number of |
| 518 | /// bytes read from `src` and written to `dest`; if the frame does not declare | 582 | /// bytes read from `src` and written to `dest`; if the frame does not declare |
| 519 | /// its decompressed content size `error.UnknownContentSizeUnsupported` is | 583 | /// its decompressed content size `error.UnknownContentSizeUnsupported` is |
| ... | @@ -521,7 +585,11 @@ const offset_table_size_max = 1 << types.compressed_block.table_accuracy_log_max | ... | @@ -521,7 +585,11 @@ const offset_table_size_max = 1 << types.compressed_block.table_accuracy_log_max |
| 521 | /// dictionary, and `error.ChecksumFailure` if `verify_checksum` is `true` and | 585 | /// dictionary, and `error.ChecksumFailure` if `verify_checksum` is `true` and |
| 522 | /// the frame contains a checksum that does not match the checksum computed from | 586 | /// the frame contains a checksum that does not match the checksum computed from |
| 523 | /// the decompressed frame. | 587 | /// the decompressed frame. |
| 524 | pub fn decodeZStandardFrame(dest: []u8, src: []const u8, verify_checksum: bool) !ReadWriteCount { | 588 | pub fn decodeZStandardFrame( |
| | 589 | dest: []u8, |
| | 590 | src: []const u8, |
| | 591 | verify_checksum: bool, |
| | 592 | ) (error{ UnknownContentSizeUnsupported, ContentTooLarge } || FrameError)!ReadWriteCount { |
| 525 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); | 593 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); |
| 526 | var consumed_count: usize = 4; | 594 | var consumed_count: usize = 4; |
| 527 | | 595 | |
| ... | @@ -530,13 +598,11 @@ pub fn decodeZStandardFrame(dest: []u8, src: []const u8, verify_checksum: bool) | ... | @@ -530,13 +598,11 @@ pub fn decodeZStandardFrame(dest: []u8, src: []const u8, verify_checksum: bool) |
| 530 | if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported; | 598 | if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported; |
| 531 | | 599 | |
| 532 | const content_size = frame_header.content_size orelse return error.UnknownContentSizeUnsupported; | 600 | const content_size = frame_header.content_size orelse return error.UnknownContentSizeUnsupported; |
| 533 | // const window_size = frameWindowSize(header) orelse return error.WindowSizeUnknown; | | |
| 534 | if (dest.len < content_size) return error.ContentTooLarge; | 601 | if (dest.len < content_size) return error.ContentTooLarge; |
| 535 | | 602 | |
| 536 | const should_compute_checksum = frame_header.descriptor.content_checksum_flag and verify_checksum; | 603 | const should_compute_checksum = frame_header.descriptor.content_checksum_flag and verify_checksum; |
| 537 | var hash_state = if (should_compute_checksum) std.hash.XxHash64.init(0) else undefined; | 604 | var hash_state = if (should_compute_checksum) std.hash.XxHash64.init(0) else undefined; |
| 538 | | 605 | |
| 539 | // TODO: block_maximum_size should be @min(1 << 17, window_size); | | |
| 540 | const written_count = try decodeFrameBlocks( | 606 | const written_count = try decodeFrameBlocks( |
| 541 | dest, | 607 | dest, |
| 542 | src[consumed_count..], | 608 | src[consumed_count..], |
| ... | @@ -567,7 +633,7 @@ pub fn decodeZStandardFrameAlloc( | ... | @@ -567,7 +633,7 @@ pub fn decodeZStandardFrameAlloc( |
| 567 | src: []const u8, | 633 | src: []const u8, |
| 568 | verify_checksum: bool, | 634 | verify_checksum: bool, |
| 569 | window_size_max: usize, | 635 | window_size_max: usize, |
| 570 | ) ![]u8 { | 636 | ) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory } || FrameError)![]u8 { |
| 571 | var result = std.ArrayList(u8).init(allocator); | 637 | var result = std.ArrayList(u8).init(allocator); |
| 572 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); | 638 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); |
| 573 | var consumed_count: usize = 4; | 639 | var consumed_count: usize = 4; |
| ... | @@ -628,7 +694,7 @@ pub fn decodeZStandardFrameAlloc( | ... | @@ -628,7 +694,7 @@ pub fn decodeZStandardFrameAlloc( |
| 628 | block_header = decodeBlockHeader(src[consumed_count..][0..3]); | 694 | block_header = decodeBlockHeader(src[consumed_count..][0..3]); |
| 629 | consumed_count += 3; | 695 | consumed_count += 3; |
| 630 | }) { | 696 | }) { |
| 631 | if (block_header.block_size > block_size_maximum) return error.CompressedBlockSizeOverMaximum; | 697 | if (block_header.block_size > block_size_maximum) return error.BlockSizeOverMaximum; |
| 632 | const written_size = try decodeBlockRingBuffer( | 698 | const written_size = try decodeBlockRingBuffer( |
| 633 | &ring_buffer, | 699 | &ring_buffer, |
| 634 | src[consumed_count..], | 700 | src[consumed_count..], |
| ... | @@ -637,7 +703,7 @@ pub fn decodeZStandardFrameAlloc( | ... | @@ -637,7 +703,7 @@ pub fn decodeZStandardFrameAlloc( |
| 637 | &consumed_count, | 703 | &consumed_count, |
| 638 | block_size_maximum, | 704 | block_size_maximum, |
| 639 | ); | 705 | ); |
| 640 | if (written_size > block_size_maximum) return error.DecompressedBlockSizeOverMaximum; | 706 | if (written_size > block_size_maximum) return error.BlockSizeOverMaximum; |
| 641 | const written_slice = ring_buffer.sliceLast(written_size); | 707 | const written_slice = ring_buffer.sliceLast(written_size); |
| 642 | try result.appendSlice(written_slice.first); | 708 | try result.appendSlice(written_slice.first); |
| 643 | try result.appendSlice(written_slice.second); | 709 | try result.appendSlice(written_slice.second); |
| ... | @@ -650,8 +716,21 @@ pub fn decodeZStandardFrameAlloc( | ... | @@ -650,8 +716,21 @@ pub fn decodeZStandardFrameAlloc( |
| 650 | return result.toOwnedSlice(); | 716 | return result.toOwnedSlice(); |
| 651 | } | 717 | } |
| 652 | | 718 | |
| | 719 | const DecodeBlockError = error{ |
| | 720 | BlockSizeOverMaximum, |
| | 721 | MalformedBlockSize, |
| | 722 | ReservedBlock, |
| | 723 | MalformedRleBlock, |
| | 724 | MalformedCompressedBlock, |
| | 725 | }; |
| | 726 | |
| 653 | /// Convenience wrapper for decoding all blocks in a frame; see `decodeBlock()`. | 727 | /// Convenience wrapper for decoding all blocks in a frame; see `decodeBlock()`. |
| 654 | pub fn decodeFrameBlocks(dest: []u8, src: []const u8, consumed_count: *usize, hash: ?*std.hash.XxHash64) !usize { | 728 | pub fn decodeFrameBlocks( |
| | 729 | dest: []u8, |
| | 730 | src: []const u8, |
| | 731 | consumed_count: *usize, |
| | 732 | hash: ?*std.hash.XxHash64, |
| | 733 | ) DecodeBlockError!usize { |
| 655 | // These tables take 7680 bytes | 734 | // These tables take 7680 bytes |
| 656 | var literal_fse_data: [literal_table_size_max]Table.Fse = undefined; | 735 | var literal_fse_data: [literal_table_size_max]Table.Fse = undefined; |
| 657 | var match_fse_data: [match_table_size_max]Table.Fse = undefined; | 736 | var match_fse_data: [match_table_size_max]Table.Fse = undefined; |
| ... | @@ -702,7 +781,12 @@ pub fn decodeFrameBlocks(dest: []u8, src: []const u8, consumed_count: *usize, ha | ... | @@ -702,7 +781,12 @@ pub fn decodeFrameBlocks(dest: []u8, src: []const u8, consumed_count: *usize, ha |
| 702 | return written_count; | 781 | return written_count; |
| 703 | } | 782 | } |
| 704 | | 783 | |
| 705 | fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) !usize { | 784 | fn decodeRawBlock( |
| | 785 | dest: []u8, |
| | 786 | src: []const u8, |
| | 787 | block_size: u21, |
| | 788 | consumed_count: *usize, |
| | 789 | ) error{MalformedBlockSize}!usize { |
| 706 | if (src.len < block_size) return error.MalformedBlockSize; | 790 | if (src.len < block_size) return error.MalformedBlockSize; |
| 707 | const data = src[0..block_size]; | 791 | const data = src[0..block_size]; |
| 708 | std.mem.copy(u8, dest, data); | 792 | std.mem.copy(u8, dest, data); |
| ... | @@ -710,7 +794,12 @@ fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: | ... | @@ -710,7 +794,12 @@ fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: |
| 710 | return block_size; | 794 | return block_size; |
| 711 | } | 795 | } |
| 712 | | 796 | |
| 713 | fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) !usize { | 797 | fn decodeRawBlockRingBuffer( |
| | 798 | dest: *RingBuffer, |
| | 799 | src: []const u8, |
| | 800 | block_size: u21, |
| | 801 | consumed_count: *usize, |
| | 802 | ) error{MalformedBlockSize}!usize { |
| 714 | if (src.len < block_size) return error.MalformedBlockSize; | 803 | if (src.len < block_size) return error.MalformedBlockSize; |
| 715 | const data = src[0..block_size]; | 804 | const data = src[0..block_size]; |
| 716 | dest.writeSliceAssumeCapacity(data); | 805 | dest.writeSliceAssumeCapacity(data); |
| ... | @@ -718,7 +807,12 @@ fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, | ... | @@ -718,7 +807,12 @@ fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, |
| 718 | return block_size; | 807 | return block_size; |
| 719 | } | 808 | } |
| 720 | | 809 | |
| 721 | fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: *usize) !usize { | 810 | fn decodeRleBlock( |
| | 811 | dest: []u8, |
| | 812 | src: []const u8, |
| | 813 | block_size: u21, |
| | 814 | consumed_count: *usize, |
| | 815 | ) error{MalformedRleBlock}!usize { |
| 722 | if (src.len < 1) return error.MalformedRleBlock; | 816 | if (src.len < 1) return error.MalformedRleBlock; |
| 723 | var write_pos: usize = 0; | 817 | var write_pos: usize = 0; |
| 724 | while (write_pos < block_size) : (write_pos += 1) { | 818 | while (write_pos < block_size) : (write_pos += 1) { |
| ... | @@ -728,7 +822,12 @@ fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: | ... | @@ -728,7 +822,12 @@ fn decodeRleBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: |
| 728 | return block_size; | 822 | return block_size; |
| 729 | } | 823 | } |
| 730 | | 824 | |
| 731 | fn decodeRleBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, consumed_count: *usize) !usize { | 825 | fn decodeRleBlockRingBuffer( |
| | 826 | dest: *RingBuffer, |
| | 827 | src: []const u8, |
| | 828 | block_size: u21, |
| | 829 | consumed_count: *usize, |
| | 830 | ) error{MalformedRleBlock}!usize { |
| 732 | if (src.len < 1) return error.MalformedRleBlock; | 831 | if (src.len < 1) return error.MalformedRleBlock; |
| 733 | var write_pos: usize = 0; | 832 | var write_pos: usize = 0; |
| 734 | while (write_pos < block_size) : (write_pos += 1) { | 833 | while (write_pos < block_size) : (write_pos += 1) { |
| ... | @@ -749,7 +848,7 @@ pub fn decodeBlock( | ... | @@ -749,7 +848,7 @@ pub fn decodeBlock( |
| 749 | decode_state: *DecodeState, | 848 | decode_state: *DecodeState, |
| 750 | consumed_count: *usize, | 849 | consumed_count: *usize, |
| 751 | written_count: usize, | 850 | written_count: usize, |
| 752 | ) !usize { | 851 | ) DecodeBlockError!usize { |
| 753 | const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB | 852 | const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB |
| 754 | const block_size = block_header.block_size; | 853 | const block_size = block_header.block_size; |
| 755 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; | 854 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; |
| ... | @@ -759,31 +858,33 @@ pub fn decodeBlock( | ... | @@ -759,31 +858,33 @@ pub fn decodeBlock( |
| 759 | .compressed => { | 858 | .compressed => { |
| 760 | if (src.len < block_size) return error.MalformedBlockSize; | 859 | if (src.len < block_size) return error.MalformedBlockSize; |
| 761 | var bytes_read: usize = 0; | 860 | var bytes_read: usize = 0; |
| 762 | const literals = try decodeLiteralsSection(src, &bytes_read); | 861 | const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock; |
| 763 | const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read); | 862 | const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch |
| | 863 | return error.MalformedCompressedBlock; |
| 764 | | 864 | |
| 765 | bytes_read += try decode_state.prepare(src[bytes_read..], literals, sequences_header); | 865 | bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch |
| | 866 | return error.MalformedCompressedBlock; |
| 766 | | 867 | |
| 767 | var bytes_written: usize = 0; | 868 | var bytes_written: usize = 0; |
| 768 | if (sequences_header.sequence_count > 0) { | 869 | if (sequences_header.sequence_count > 0) { |
| 769 | const bit_stream_bytes = src[bytes_read..block_size]; | 870 | const bit_stream_bytes = src[bytes_read..block_size]; |
| 770 | var bit_stream: ReverseBitReader = undefined; | 871 | var bit_stream: ReverseBitReader = undefined; |
| 771 | try bit_stream.init(bit_stream_bytes); | 872 | bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock; |
| 772 | | 873 | |
| 773 | try decode_state.readInitialFseState(&bit_stream); | 874 | decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock; |
| 774 | | 875 | |
| 775 | var sequence_size_limit = block_size_max; | 876 | var sequence_size_limit = block_size_max; |
| 776 | var i: usize = 0; | 877 | var i: usize = 0; |
| 777 | while (i < sequences_header.sequence_count) : (i += 1) { | 878 | while (i < sequences_header.sequence_count) : (i += 1) { |
| 778 | const write_pos = written_count + bytes_written; | 879 | const write_pos = written_count + bytes_written; |
| 779 | const decompressed_size = try decode_state.decodeSequenceSlice( | 880 | const decompressed_size = decode_state.decodeSequenceSlice( |
| 780 | dest, | 881 | dest, |
| 781 | write_pos, | 882 | write_pos, |
| 782 | literals, | 883 | literals, |
| 783 | &bit_stream, | 884 | &bit_stream, |
| 784 | sequence_size_limit, | 885 | sequence_size_limit, |
| 785 | i == sequences_header.sequence_count - 1, | 886 | i == sequences_header.sequence_count - 1, |
| 786 | ); | 887 | ) catch return error.MalformedCompressedBlock; |
| 787 | bytes_written += decompressed_size; | 888 | bytes_written += decompressed_size; |
| 788 | sequence_size_limit -= decompressed_size; | 889 | sequence_size_limit -= decompressed_size; |
| 789 | } | 890 | } |
| ... | @@ -793,7 +894,8 @@ pub fn decodeBlock( | ... | @@ -793,7 +894,8 @@ pub fn decodeBlock( |
| 793 | | 894 | |
| 794 | if (decode_state.literal_written_count < literals.header.regenerated_size) { | 895 | if (decode_state.literal_written_count < literals.header.regenerated_size) { |
| 795 | const len = literals.header.regenerated_size - decode_state.literal_written_count; | 896 | const len = literals.header.regenerated_size - decode_state.literal_written_count; |
| 796 | try decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], literals, len); | 897 | decode_state.decodeLiteralsSlice(dest[written_count + bytes_written ..], literals, len) catch |
| | 898 | return error.MalformedCompressedBlock; |
| 797 | bytes_written += len; | 899 | bytes_written += len; |
| 798 | } | 900 | } |
| 799 | | 901 | |
| ... | @@ -802,7 +904,7 @@ pub fn decodeBlock( | ... | @@ -802,7 +904,7 @@ pub fn decodeBlock( |
| 802 | consumed_count.* += bytes_read; | 904 | consumed_count.* += bytes_read; |
| 803 | return bytes_written; | 905 | return bytes_written; |
| 804 | }, | 906 | }, |
| 805 | .reserved => return error.FrameContainsReservedBlock, | 907 | .reserved => return error.ReservedBlock, |
| 806 | } | 908 | } |
| 807 | } | 909 | } |
| 808 | | 910 | |
| ... | @@ -816,7 +918,7 @@ pub fn decodeBlockRingBuffer( | ... | @@ -816,7 +918,7 @@ pub fn decodeBlockRingBuffer( |
| 816 | decode_state: *DecodeState, | 918 | decode_state: *DecodeState, |
| 817 | consumed_count: *usize, | 919 | consumed_count: *usize, |
| 818 | block_size_max: usize, | 920 | block_size_max: usize, |
| 819 | ) !usize { | 921 | ) DecodeBlockError!usize { |
| 820 | const block_size = block_header.block_size; | 922 | const block_size = block_header.block_size; |
| 821 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; | 923 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; |
| 822 | switch (block_header.block_type) { | 924 | switch (block_header.block_type) { |
| ... | @@ -825,29 +927,31 @@ pub fn decodeBlockRingBuffer( | ... | @@ -825,29 +927,31 @@ pub fn decodeBlockRingBuffer( |
| 825 | .compressed => { | 927 | .compressed => { |
| 826 | if (src.len < block_size) return error.MalformedBlockSize; | 928 | if (src.len < block_size) return error.MalformedBlockSize; |
| 827 | var bytes_read: usize = 0; | 929 | var bytes_read: usize = 0; |
| 828 | const literals = try decodeLiteralsSection(src, &bytes_read); | 930 | const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock; |
| 829 | const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read); | 931 | const sequences_header = decodeSequencesHeader(src[bytes_read..], &bytes_read) catch |
| | 932 | return error.MalformedCompressedBlock; |
| 830 | | 933 | |
| 831 | bytes_read += try decode_state.prepare(src[bytes_read..], literals, sequences_header); | 934 | bytes_read += decode_state.prepare(src[bytes_read..], literals, sequences_header) catch |
| | 935 | return error.MalformedCompressedBlock; |
| 832 | | 936 | |
| 833 | var bytes_written: usize = 0; | 937 | var bytes_written: usize = 0; |
| 834 | if (sequences_header.sequence_count > 0) { | 938 | if (sequences_header.sequence_count > 0) { |
| 835 | const bit_stream_bytes = src[bytes_read..block_size]; | 939 | const bit_stream_bytes = src[bytes_read..block_size]; |
| 836 | var bit_stream: ReverseBitReader = undefined; | 940 | var bit_stream: ReverseBitReader = undefined; |
| 837 | try bit_stream.init(bit_stream_bytes); | 941 | bit_stream.init(bit_stream_bytes) catch return error.MalformedCompressedBlock; |
| 838 | | 942 | |
| 839 | try decode_state.readInitialFseState(&bit_stream); | 943 | decode_state.readInitialFseState(&bit_stream) catch return error.MalformedCompressedBlock; |
| 840 | | 944 | |
| 841 | var sequence_size_limit = block_size_max; | 945 | var sequence_size_limit = block_size_max; |
| 842 | var i: usize = 0; | 946 | var i: usize = 0; |
| 843 | while (i < sequences_header.sequence_count) : (i += 1) { | 947 | while (i < sequences_header.sequence_count) : (i += 1) { |
| 844 | const decompressed_size = try decode_state.decodeSequenceRingBuffer( | 948 | const decompressed_size = decode_state.decodeSequenceRingBuffer( |
| 845 | dest, | 949 | dest, |
| 846 | literals, | 950 | literals, |
| 847 | &bit_stream, | 951 | &bit_stream, |
| 848 | sequence_size_limit, | 952 | sequence_size_limit, |
| 849 | i == sequences_header.sequence_count - 1, | 953 | i == sequences_header.sequence_count - 1, |
| 850 | ); | 954 | ) catch return error.MalformedCompressedBlock; |
| 851 | bytes_written += decompressed_size; | 955 | bytes_written += decompressed_size; |
| 852 | sequence_size_limit -= decompressed_size; | 956 | sequence_size_limit -= decompressed_size; |
| 853 | } | 957 | } |
| ... | @@ -857,7 +961,8 @@ pub fn decodeBlockRingBuffer( | ... | @@ -857,7 +961,8 @@ pub fn decodeBlockRingBuffer( |
| 857 | | 961 | |
| 858 | if (decode_state.literal_written_count < literals.header.regenerated_size) { | 962 | if (decode_state.literal_written_count < literals.header.regenerated_size) { |
| 859 | const len = literals.header.regenerated_size - decode_state.literal_written_count; | 963 | const len = literals.header.regenerated_size - decode_state.literal_written_count; |
| 860 | try decode_state.decodeLiteralsRingBuffer(dest, literals, len); | 964 | decode_state.decodeLiteralsRingBuffer(dest, literals, len) catch |
| | 965 | return error.MalformedCompressedBlock; |
| 861 | bytes_written += len; | 966 | bytes_written += len; |
| 862 | } | 967 | } |
| 863 | | 968 | |
| ... | @@ -866,7 +971,7 @@ pub fn decodeBlockRingBuffer( | ... | @@ -866,7 +971,7 @@ pub fn decodeBlockRingBuffer( |
| 866 | consumed_count.* += bytes_read; | 971 | consumed_count.* += bytes_read; |
| 867 | return bytes_written; | 972 | return bytes_written; |
| 868 | }, | 973 | }, |
| 869 | .reserved => return error.FrameContainsReservedBlock, | 974 | .reserved => return error.ReservedBlock, |
| 870 | } | 975 | } |
| 871 | } | 976 | } |
| 872 | | 977 | |
| ... | @@ -901,9 +1006,10 @@ pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 { | ... | @@ -901,9 +1006,10 @@ pub fn frameWindowSize(header: frame.ZStandard.Header) ?u64 { |
| 901 | } else return header.content_size; | 1006 | } else return header.content_size; |
| 902 | } | 1007 | } |
| 903 | | 1008 | |
| | 1009 | const InvalidBit = error{ UnusedBitSet, ReservedBitSet }; |
| 904 | /// Decode the header of a Zstandard frame. Returns `error.UnusedBitSet` or | 1010 | /// Decode the header of a Zstandard frame. Returns `error.UnusedBitSet` or |
| 905 | /// `error.ReservedBitSet` if the corresponding bits are sets. | 1011 | /// `error.ReservedBitSet` if the corresponding bits are sets. |
| 906 | pub fn decodeZStandardHeader(src: []const u8, consumed_count: ?*usize) !frame.ZStandard.Header { | 1012 | pub fn decodeZStandardHeader(src: []const u8, consumed_count: ?*usize) InvalidBit!frame.ZStandard.Header { |
| 907 | const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, src[0]); | 1013 | const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, src[0]); |
| 908 | | 1014 | |
| 909 | if (descriptor.unused) return error.UnusedBitSet; | 1015 | if (descriptor.unused) return error.UnusedBitSet; |
| ... | @@ -958,7 +1064,10 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header { | ... | @@ -958,7 +1064,10 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header { |
| 958 | | 1064 | |
| 959 | /// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the | 1065 | /// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the |
| 960 | /// number of bytes the section uses. | 1066 | /// number of bytes the section uses. |
| 961 | pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsSection { | 1067 | pub fn decodeLiteralsSection( |
| | 1068 | src: []const u8, |
| | 1069 | consumed_count: *usize, |
| | 1070 | ) (error{ MalformedLiteralsHeader, MalformedLiteralsSection } || DecodeHuffmanError)!LiteralsSection { |
| 962 | var bytes_read: usize = 0; | 1071 | var bytes_read: usize = 0; |
| 963 | const header = try decodeLiteralsHeader(src, &bytes_read); | 1072 | const header = try decodeLiteralsHeader(src, &bytes_read); |
| 964 | switch (header.block_type) { | 1073 | switch (header.block_type) { |
| ... | @@ -1032,7 +1141,13 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS | ... | @@ -1032,7 +1141,13 @@ pub fn decodeLiteralsSection(src: []const u8, consumed_count: *usize) !LiteralsS |
| 1032 | } | 1141 | } |
| 1033 | } | 1142 | } |
| 1034 | | 1143 | |
| 1035 | fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.HuffmanTree { | 1144 | const DecodeHuffmanError = error{ |
| | 1145 | MalformedHuffmanTree, |
| | 1146 | MalformedFseTable, |
| | 1147 | MalformedAccuracyLog, |
| | 1148 | }; |
| | 1149 | |
| | 1150 | fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) DecodeHuffmanError!LiteralsSection.HuffmanTree { |
| 1036 | var bytes_read: usize = 0; | 1151 | var bytes_read: usize = 0; |
| 1037 | bytes_read += 1; | 1152 | bytes_read += 1; |
| 1038 | if (src.len == 0) return error.MalformedHuffmanTree; | 1153 | if (src.len == 0) return error.MalformedHuffmanTree; |
| ... | @@ -1049,22 +1164,25 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H | ... | @@ -1049,22 +1164,25 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H |
| 1049 | var bit_reader = bitReader(counting_reader.reader()); | 1164 | var bit_reader = bitReader(counting_reader.reader()); |
| 1050 | | 1165 | |
| 1051 | var entries: [1 << 6]Table.Fse = undefined; | 1166 | var entries: [1 << 6]Table.Fse = undefined; |
| 1052 | const table_size = try decodeFseTable(&bit_reader, 256, 6, &entries); | 1167 | const table_size = decodeFseTable(&bit_reader, 256, 6, &entries) catch |err| switch (err) { |
| | 1168 | error.MalformedAccuracyLog, error.MalformedFseTable => |e| return e, |
| | 1169 | error.EndOfStream => return error.MalformedFseTable, |
| | 1170 | }; |
| 1053 | const accuracy_log = std.math.log2_int_ceil(usize, table_size); | 1171 | const accuracy_log = std.math.log2_int_ceil(usize, table_size); |
| 1054 | | 1172 | |
| 1055 | const start_index = std.math.cast(usize, 1 + counting_reader.bytes_read) orelse return error.MalformedHuffmanTree; | 1173 | const start_index = std.math.cast(usize, 1 + counting_reader.bytes_read) orelse return error.MalformedHuffmanTree; |
| 1056 | var huff_data = src[start_index .. compressed_size + 1]; | 1174 | var huff_data = src[start_index .. compressed_size + 1]; |
| 1057 | var huff_bits: ReverseBitReader = undefined; | 1175 | var huff_bits: ReverseBitReader = undefined; |
| 1058 | try huff_bits.init(huff_data); | 1176 | huff_bits.init(huff_data) catch return error.MalformedHuffmanTree; |
| 1059 | | 1177 | |
| 1060 | var i: usize = 0; | 1178 | var i: usize = 0; |
| 1061 | var even_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log); | 1179 | var even_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree; |
| 1062 | var odd_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log); | 1180 | var odd_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree; |
| 1063 | | 1181 | |
| 1064 | while (i < 255) { | 1182 | while (i < 255) { |
| 1065 | const even_data = entries[even_state]; | 1183 | const even_data = entries[even_state]; |
| 1066 | var read_bits: usize = 0; | 1184 | var read_bits: usize = 0; |
| 1067 | const even_bits = try huff_bits.readBits(u32, even_data.bits, &read_bits); | 1185 | const even_bits = huff_bits.readBits(u32, even_data.bits, &read_bits) catch unreachable; |
| 1068 | weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree; | 1186 | weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree; |
| 1069 | i += 1; | 1187 | i += 1; |
| 1070 | if (read_bits < even_data.bits) { | 1188 | if (read_bits < even_data.bits) { |
| ... | @@ -1076,7 +1194,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H | ... | @@ -1076,7 +1194,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H |
| 1076 | | 1194 | |
| 1077 | read_bits = 0; | 1195 | read_bits = 0; |
| 1078 | const odd_data = entries[odd_state]; | 1196 | const odd_data = entries[odd_state]; |
| 1079 | const odd_bits = try huff_bits.readBits(u32, odd_data.bits, &read_bits); | 1197 | const odd_bits = huff_bits.readBits(u32, odd_data.bits, &read_bits) catch unreachable; |
| 1080 | weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree; | 1198 | weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree; |
| 1081 | i += 1; | 1199 | i += 1; |
| 1082 | if (read_bits < odd_data.bits) { | 1200 | if (read_bits < odd_data.bits) { |
| ... | @@ -1177,8 +1295,8 @@ fn lessThanByWeight( | ... | @@ -1177,8 +1295,8 @@ fn lessThanByWeight( |
| 1177 | } | 1295 | } |
| 1178 | | 1296 | |
| 1179 | /// Decode a literals section header. | 1297 | /// Decode a literals section header. |
| 1180 | pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSection.Header { | 1298 | pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{MalformedLiteralsHeader}!LiteralsSection.Header { |
| 1181 | if (src.len == 0) return error.MalformedLiteralsSection; | 1299 | if (src.len == 0) return error.MalformedLiteralsHeader; |
| 1182 | const byte0 = src[0]; | 1300 | const byte0 = src[0]; |
| 1183 | const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11); | 1301 | const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11); |
| 1184 | const size_format = @intCast(u2, (byte0 & 0b1100) >> 2); | 1302 | const size_format = @intCast(u2, (byte0 & 0b1100) >> 2); |
| ... | @@ -1243,8 +1361,11 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSe | ... | @@ -1243,8 +1361,11 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSe |
| 1243 | } | 1361 | } |
| 1244 | | 1362 | |
| 1245 | /// Decode a sequences section header. | 1363 | /// Decode a sequences section header. |
| 1246 | pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !SequencesSection.Header { | 1364 | pub fn decodeSequencesHeader( |
| 1247 | if (src.len == 0) return error.MalformedSequencesSection; | 1365 | src: []const u8, |
| | 1366 | consumed_count: *usize, |
| | 1367 | ) error{ MalformedSequencesHeader, ReservedBitSet }!SequencesSection.Header { |
| | 1368 | if (src.len == 0) return error.MalformedSequencesHeader; |
| 1248 | var sequence_count: u24 = undefined; | 1369 | var sequence_count: u24 = undefined; |
| 1249 | | 1370 | |
| 1250 | var bytes_read: usize = 0; | 1371 | var bytes_read: usize = 0; |
| ... | @@ -1262,16 +1383,16 @@ pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !Sequences | ... | @@ -1262,16 +1383,16 @@ pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !Sequences |
| 1262 | sequence_count = byte0; | 1383 | sequence_count = byte0; |
| 1263 | bytes_read += 1; | 1384 | bytes_read += 1; |
| 1264 | } else if (byte0 < 255) { | 1385 | } else if (byte0 < 255) { |
| 1265 | if (src.len < 2) return error.MalformedSequencesSection; | 1386 | if (src.len < 2) return error.MalformedSequencesHeader; |
| 1266 | sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1]; | 1387 | sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1]; |
| 1267 | bytes_read += 2; | 1388 | bytes_read += 2; |
| 1268 | } else { | 1389 | } else { |
| 1269 | if (src.len < 3) return error.MalformedSequencesSection; | 1390 | if (src.len < 3) return error.MalformedSequencesHeader; |
| 1270 | sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00; | 1391 | sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00; |
| 1271 | bytes_read += 3; | 1392 | bytes_read += 3; |
| 1272 | } | 1393 | } |
| 1273 | | 1394 | |
| 1274 | if (src.len < bytes_read + 1) return error.MalformedSequencesSection; | 1395 | if (src.len < bytes_read + 1) return error.MalformedSequencesHeader; |
| 1275 | const compression_modes = src[bytes_read]; | 1396 | const compression_modes = src[bytes_read]; |
| 1276 | bytes_read += 1; | 1397 | bytes_read += 1; |
| 1277 | | 1398 | |
| ... | @@ -1441,17 +1562,17 @@ pub const ReverseBitReader = struct { | ... | @@ -1441,17 +1562,17 @@ pub const ReverseBitReader = struct { |
| 1441 | byte_reader: ReversedByteReader, | 1562 | byte_reader: ReversedByteReader, |
| 1442 | bit_reader: std.io.BitReader(.Big, ReversedByteReader.Reader), | 1563 | bit_reader: std.io.BitReader(.Big, ReversedByteReader.Reader), |
| 1443 | | 1564 | |
| 1444 | pub fn init(self: *ReverseBitReader, bytes: []const u8) !void { | 1565 | pub fn init(self: *ReverseBitReader, bytes: []const u8) error{BitStreamHasNoStartBit}!void { |
| 1445 | self.byte_reader = ReversedByteReader.init(bytes); | 1566 | self.byte_reader = ReversedByteReader.init(bytes); |
| 1446 | self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader()); | 1567 | self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader()); |
| 1447 | while (0 == self.readBitsNoEof(u1, 1) catch return error.BitStreamHasNoStartBit) {} | 1568 | while (0 == self.readBitsNoEof(u1, 1) catch return error.BitStreamHasNoStartBit) {} |
| 1448 | } | 1569 | } |
| 1449 | | 1570 | |
| 1450 | pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) !U { | 1571 | pub fn readBitsNoEof(self: *@This(), comptime U: type, num_bits: usize) error{EndOfStream}!U { |
| 1451 | return self.bit_reader.readBitsNoEof(U, num_bits); | 1572 | return self.bit_reader.readBitsNoEof(U, num_bits); |
| 1452 | } | 1573 | } |
| 1453 | | 1574 | |
| 1454 | pub fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) !U { | 1575 | pub fn readBits(self: *@This(), comptime U: type, num_bits: usize, out_bits: *usize) error{}!U { |
| 1455 | return try self.bit_reader.readBits(U, num_bits, out_bits); | 1576 | return try self.bit_reader.readBits(U, num_bits, out_bits); |
| 1456 | } | 1577 | } |
| 1457 | | 1578 | |