| ... | ... | @@ -22,7 +22,7 @@ fn isSkippableMagic(magic: u32) bool { |
| 22 | 22 | /// if the the frame is skippable, `null` for Zstanndard frames that do not |
| 23 | 23 | /// declare their content size. Returns `UnusedBitSet` and `ReservedBitSet` |
| 24 | 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 | 26 | switch (try frameType(src)) { |
| 27 | 27 | .zstandard => { |
| 28 | 28 | const header = try decodeZStandardHeader(src[4..], null); |
| ... | ... | @@ -52,7 +52,11 @@ const ReadWriteCount = struct { |
| 52 | 52 | |
| 53 | 53 | /// Decodes the frame at the start of `src` into `dest`. Returns the number of |
| 54 | 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 | 60 | return switch (try frameType(src)) { |
| 57 | 61 | .zstandard => decodeZStandardFrame(dest, src, verify_checksum), |
| 58 | 62 | .skippable => ReadWriteCount{ |
| ... | ... | @@ -100,7 +104,7 @@ pub const DecodeState = struct { |
| 100 | 104 | src: []const u8, |
| 101 | 105 | literals: LiteralsSection, |
| 102 | 106 | sequences_header: SequencesSection.Header, |
| 103 | | ) !usize { |
| 107 | ) (error{ BitStreamHasNoStartBit, TreelessLiteralsFirst } || FseTableError)!usize { |
| 104 | 108 | if (literals.huffman_tree) |tree| { |
| 105 | 109 | self.huffman_tree = tree; |
| 106 | 110 | } else if (literals.header.block_type == .treeless and self.huffman_tree == null) { |
| ... | ... | @@ -145,7 +149,7 @@ pub const DecodeState = struct { |
| 145 | 149 | |
| 146 | 150 | /// Read initial FSE states for sequence decoding. Returns `error.EndOfStream` |
| 147 | 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 | 153 | self.literal.state = try bit_reader.readBitsNoEof(u9, self.literal.accuracy_log); |
| 150 | 154 | self.offset.state = try bit_reader.readBitsNoEof(u8, self.offset.accuracy_log); |
| 151 | 155 | self.match.state = try bit_reader.readBitsNoEof(u9, self.match.accuracy_log); |
| ... | ... | @@ -169,7 +173,11 @@ pub const DecodeState = struct { |
| 169 | 173 | |
| 170 | 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 | 181 | switch (@field(self, @tagName(choice)).table) { |
| 174 | 182 | .rle => {}, |
| 175 | 183 | .fse => |table| { |
| ... | ... | @@ -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 | 203 | fn updateFseTable( |
| 189 | 204 | self: *DecodeState, |
| 190 | 205 | src: []const u8, |
| 191 | 206 | comptime choice: DataType, |
| 192 | 207 | mode: SequencesSection.Header.Mode, |
| 193 | | ) !usize { |
| 208 | ) FseTableError!usize { |
| 194 | 209 | const field_name = @tagName(choice); |
| 195 | 210 | switch (mode) { |
| 196 | 211 | .predefined => { |
| 197 | | @field(self, field_name).accuracy_log = @field(types.compressed_block.default_accuracy_log, field_name); |
| 198 | | @field(self, field_name).table = @field(types.compressed_block, "predefined_" ++ field_name ++ "_fse_table"); |
| 212 | @field(self, field_name).accuracy_log = |
| 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 | 217 | return 0; |
| 200 | 218 | }, |
| 201 | 219 | .rle => { |
| ... | ... | @@ -214,9 +232,11 @@ pub const DecodeState = struct { |
| 214 | 232 | @field(types.compressed_block.table_accuracy_log_max, field_name), |
| 215 | 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 | 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 | 241 | .repeat => return if (self.fse_tables_undefined) error.RepeatModeFirst else 0, |
| 222 | 242 | } |
| ... | ... | @@ -228,7 +248,10 @@ pub const DecodeState = struct { |
| 228 | 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 | 255 | const raw_code = self.getCode(.offset); |
| 233 | 256 | const offset_code = std.math.cast(u5, raw_code) orelse { |
| 234 | 257 | return error.OffsetCodeTooLarge; |
| ... | ... | @@ -272,7 +295,7 @@ pub const DecodeState = struct { |
| 272 | 295 | write_pos: usize, |
| 273 | 296 | literals: LiteralsSection, |
| 274 | 297 | sequence: Sequence, |
| 275 | | ) !void { |
| 298 | ) (error{MalformedSequence} || DecodeLiteralsError)!void { |
| 276 | 299 | if (sequence.offset > write_pos + sequence.literal_length) return error.MalformedSequence; |
| 277 | 300 | |
| 278 | 301 | try self.decodeLiteralsSlice(dest[write_pos..], literals, sequence.literal_length); |
| ... | ... | @@ -288,16 +311,23 @@ pub const DecodeState = struct { |
| 288 | 311 | dest: *RingBuffer, |
| 289 | 312 | literals: LiteralsSection, |
| 290 | 313 | sequence: Sequence, |
| 291 | | ) !void { |
| 314 | ) (error{MalformedSequence} || DecodeLiteralsError)!void { |
| 292 | 315 | if (sequence.offset > dest.data.len) return error.MalformedSequence; |
| 293 | 316 | |
| 294 | 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 | 320 | // TODO: would std.mem.copy and figuring out dest slice be better/faster? |
| 297 | 321 | for (copy_slice.first) |b| dest.writeAssumeCapacity(b); |
| 298 | 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 | 331 | /// Decode one sequence from `bit_reader` into `dest`, written starting at |
| 302 | 332 | /// `write_pos` and update FSE states if `last_sequence` is `false`. Returns |
| 303 | 333 | /// `error.MalformedSequence` error if the decompressed sequence would be longer |
| ... | ... | @@ -311,10 +341,10 @@ pub const DecodeState = struct { |
| 311 | 341 | dest: []u8, |
| 312 | 342 | write_pos: usize, |
| 313 | 343 | literals: LiteralsSection, |
| 314 | | bit_reader: anytype, |
| 344 | bit_reader: *ReverseBitReader, |
| 315 | 345 | sequence_size_limit: usize, |
| 316 | 346 | last_sequence: bool, |
| 317 | | ) !usize { |
| 347 | ) DecodeSequenceError!usize { |
| 318 | 348 | const sequence = try self.nextSequence(bit_reader); |
| 319 | 349 | const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length; |
| 320 | 350 | if (sequence_length > sequence_size_limit) return error.MalformedSequence; |
| ... | ... | @@ -336,7 +366,7 @@ pub const DecodeState = struct { |
| 336 | 366 | bit_reader: anytype, |
| 337 | 367 | sequence_size_limit: usize, |
| 338 | 368 | last_sequence: bool, |
| 339 | | ) !usize { |
| 369 | ) DecodeSequenceError!usize { |
| 340 | 370 | const sequence = try self.nextSequence(bit_reader); |
| 341 | 371 | const sequence_length = @as(usize, sequence.literal_length) + sequence.match_length; |
| 342 | 372 | if (sequence_length > sequence_size_limit) return error.MalformedSequence; |
| ... | ... | @@ -350,26 +380,63 @@ pub const DecodeState = struct { |
| 350 | 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 | 387 | self.literal_stream_index += 1; |
| 355 | 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 | 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 | 421 | /// Decode `len` bytes of literals into `dest`. `literals` should be the |
| 363 | 422 | /// `LiteralsSection` that was passed to `prepare()`. Returns |
| 364 | 423 | /// `error.MalformedLiteralsLength` if the number of literal bytes decoded by |
| 365 | 424 | /// `self` plus `len` is greater than the regenerated size of `literals`. |
| 366 | 425 | /// Returns `error.UnexpectedEndOfLiteralStream` and `error.PrefixNotFound` if |
| 367 | 426 | /// there are problems decoding Huffman compressed literals. |
| 368 | | pub fn decodeLiteralsSlice(self: *DecodeState, dest: []u8, literals: LiteralsSection, len: usize) !void { |
| 369 | | if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength; |
| 427 | pub fn decodeLiteralsSlice( |
| 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 | 436 | switch (literals.header.block_type) { |
| 371 | 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 | 440 | std.mem.copy(u8, dest, literal_data); |
| 374 | 441 | self.literal_written_count += len; |
| 375 | 442 | }, |
| ... | ... | @@ -395,15 +462,7 @@ pub const DecodeState = struct { |
| 395 | 462 | while (i < len) : (i += 1) { |
| 396 | 463 | var prefix: u16 = 0; |
| 397 | 464 | while (true) { |
| 398 | | const new_bits = self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |err| |
| 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 | | }; |
| 465 | const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals); |
| 407 | 466 | prefix <<= bit_count_to_read; |
| 408 | 467 | prefix |= new_bits; |
| 409 | 468 | bits_read += bit_count_to_read; |
| ... | ... | @@ -434,11 +493,19 @@ pub const DecodeState = struct { |
| 434 | 493 | } |
| 435 | 494 | |
| 436 | 495 | /// Decode literals into `dest`; see `decodeLiteralsSlice()`. |
| 437 | | pub fn decodeLiteralsRingBuffer(self: *DecodeState, dest: *RingBuffer, literals: LiteralsSection, len: usize) !void { |
| 438 | | if (self.literal_written_count + len > literals.header.regenerated_size) return error.MalformedLiteralsLength; |
| 496 | pub fn decodeLiteralsRingBuffer( |
| 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 | 505 | switch (literals.header.block_type) { |
| 440 | 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 | 509 | dest.writeSliceAssumeCapacity(literal_data); |
| 443 | 510 | self.literal_written_count += len; |
| 444 | 511 | }, |
| ... | ... | @@ -464,15 +531,7 @@ pub const DecodeState = struct { |
| 464 | 531 | while (i < len) : (i += 1) { |
| 465 | 532 | var prefix: u16 = 0; |
| 466 | 533 | while (true) { |
| 467 | | const new_bits = self.literal_stream_reader.readBitsNoEof(u16, bit_count_to_read) catch |err| |
| 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 | | }; |
| 534 | const new_bits = try self.readLiteralsBits(u16, bit_count_to_read, literals); |
| 476 | 535 | prefix <<= bit_count_to_read; |
| 477 | 536 | prefix |= new_bits; |
| 478 | 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 | 573 | const match_table_size_max = 1 << types.compressed_block.table_accuracy_log_max.match; |
| 515 | 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 | 581 | /// Decode a Zstandard frame from `src` into `dest`, returning the number of |
| 518 | 582 | /// bytes read from `src` and written to `dest`; if the frame does not declare |
| 519 | 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 | 585 | /// dictionary, and `error.ChecksumFailure` if `verify_checksum` is `true` and |
| 522 | 586 | /// the frame contains a checksum that does not match the checksum computed from |
| 523 | 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 | 593 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); |
| 526 | 594 | var consumed_count: usize = 4; |
| 527 | 595 | |
| ... | ... | @@ -530,13 +598,11 @@ pub fn decodeZStandardFrame(dest: []u8, src: []const u8, verify_checksum: bool) |
| 530 | 598 | if (frame_header.descriptor.dictionary_id_flag != 0) return error.DictionaryIdFlagUnsupported; |
| 531 | 599 | |
| 532 | 600 | const content_size = frame_header.content_size orelse return error.UnknownContentSizeUnsupported; |
| 533 | | // const window_size = frameWindowSize(header) orelse return error.WindowSizeUnknown; |
| 534 | 601 | if (dest.len < content_size) return error.ContentTooLarge; |
| 535 | 602 | |
| 536 | 603 | const should_compute_checksum = frame_header.descriptor.content_checksum_flag and verify_checksum; |
| 537 | 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 | 606 | const written_count = try decodeFrameBlocks( |
| 541 | 607 | dest, |
| 542 | 608 | src[consumed_count..], |
| ... | ... | @@ -567,7 +633,7 @@ pub fn decodeZStandardFrameAlloc( |
| 567 | 633 | src: []const u8, |
| 568 | 634 | verify_checksum: bool, |
| 569 | 635 | window_size_max: usize, |
| 570 | | ) ![]u8 { |
| 636 | ) (error{ WindowSizeUnknown, WindowTooLarge, OutOfMemory } || FrameError)![]u8 { |
| 571 | 637 | var result = std.ArrayList(u8).init(allocator); |
| 572 | 638 | assert(readInt(u32, src[0..4]) == frame.ZStandard.magic_number); |
| 573 | 639 | var consumed_count: usize = 4; |
| ... | ... | @@ -628,7 +694,7 @@ pub fn decodeZStandardFrameAlloc( |
| 628 | 694 | block_header = decodeBlockHeader(src[consumed_count..][0..3]); |
| 629 | 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 | 698 | const written_size = try decodeBlockRingBuffer( |
| 633 | 699 | &ring_buffer, |
| 634 | 700 | src[consumed_count..], |
| ... | ... | @@ -637,7 +703,7 @@ pub fn decodeZStandardFrameAlloc( |
| 637 | 703 | &consumed_count, |
| 638 | 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 | 707 | const written_slice = ring_buffer.sliceLast(written_size); |
| 642 | 708 | try result.appendSlice(written_slice.first); |
| 643 | 709 | try result.appendSlice(written_slice.second); |
| ... | ... | @@ -650,8 +716,21 @@ pub fn decodeZStandardFrameAlloc( |
| 650 | 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 | 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 | 734 | // These tables take 7680 bytes |
| 656 | 735 | var literal_fse_data: [literal_table_size_max]Table.Fse = undefined; |
| 657 | 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 | 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 | 790 | if (src.len < block_size) return error.MalformedBlockSize; |
| 707 | 791 | const data = src[0..block_size]; |
| 708 | 792 | std.mem.copy(u8, dest, data); |
| ... | ... | @@ -710,7 +794,12 @@ fn decodeRawBlock(dest: []u8, src: []const u8, block_size: u21, consumed_count: |
| 710 | 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 | 803 | if (src.len < block_size) return error.MalformedBlockSize; |
| 715 | 804 | const data = src[0..block_size]; |
| 716 | 805 | dest.writeSliceAssumeCapacity(data); |
| ... | ... | @@ -718,7 +807,12 @@ fn decodeRawBlockRingBuffer(dest: *RingBuffer, src: []const u8, block_size: u21, |
| 718 | 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 | 816 | if (src.len < 1) return error.MalformedRleBlock; |
| 723 | 817 | var write_pos: usize = 0; |
| 724 | 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 | 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 | 831 | if (src.len < 1) return error.MalformedRleBlock; |
| 733 | 832 | var write_pos: usize = 0; |
| 734 | 833 | while (write_pos < block_size) : (write_pos += 1) { |
| ... | ... | @@ -749,7 +848,7 @@ pub fn decodeBlock( |
| 749 | 848 | decode_state: *DecodeState, |
| 750 | 849 | consumed_count: *usize, |
| 751 | 850 | written_count: usize, |
| 752 | | ) !usize { |
| 851 | ) DecodeBlockError!usize { |
| 753 | 852 | const block_size_max = @min(1 << 17, dest[written_count..].len); // 128KiB |
| 754 | 853 | const block_size = block_header.block_size; |
| 755 | 854 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; |
| ... | ... | @@ -759,31 +858,33 @@ pub fn decodeBlock( |
| 759 | 858 | .compressed => { |
| 760 | 859 | if (src.len < block_size) return error.MalformedBlockSize; |
| 761 | 860 | var bytes_read: usize = 0; |
| 762 | | const literals = try decodeLiteralsSection(src, &bytes_read); |
| 763 | | const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read); |
| 861 | const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock; |
| 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 | 868 | var bytes_written: usize = 0; |
| 768 | 869 | if (sequences_header.sequence_count > 0) { |
| 769 | 870 | const bit_stream_bytes = src[bytes_read..block_size]; |
| 770 | 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 | 876 | var sequence_size_limit = block_size_max; |
| 776 | 877 | var i: usize = 0; |
| 777 | 878 | while (i < sequences_header.sequence_count) : (i += 1) { |
| 778 | 879 | const write_pos = written_count + bytes_written; |
| 779 | | const decompressed_size = try decode_state.decodeSequenceSlice( |
| 880 | const decompressed_size = decode_state.decodeSequenceSlice( |
| 780 | 881 | dest, |
| 781 | 882 | write_pos, |
| 782 | 883 | literals, |
| 783 | 884 | &bit_stream, |
| 784 | 885 | sequence_size_limit, |
| 785 | 886 | i == sequences_header.sequence_count - 1, |
| 786 | | ); |
| 887 | ) catch return error.MalformedCompressedBlock; |
| 787 | 888 | bytes_written += decompressed_size; |
| 788 | 889 | sequence_size_limit -= decompressed_size; |
| 789 | 890 | } |
| ... | ... | @@ -793,7 +894,8 @@ pub fn decodeBlock( |
| 793 | 894 | |
| 794 | 895 | if (decode_state.literal_written_count < literals.header.regenerated_size) { |
| 795 | 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 | 899 | bytes_written += len; |
| 798 | 900 | } |
| 799 | 901 | |
| ... | ... | @@ -802,7 +904,7 @@ pub fn decodeBlock( |
| 802 | 904 | consumed_count.* += bytes_read; |
| 803 | 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 | 918 | decode_state: *DecodeState, |
| 817 | 919 | consumed_count: *usize, |
| 818 | 920 | block_size_max: usize, |
| 819 | | ) !usize { |
| 921 | ) DecodeBlockError!usize { |
| 820 | 922 | const block_size = block_header.block_size; |
| 821 | 923 | if (block_size_max < block_size) return error.BlockSizeOverMaximum; |
| 822 | 924 | switch (block_header.block_type) { |
| ... | ... | @@ -825,29 +927,31 @@ pub fn decodeBlockRingBuffer( |
| 825 | 927 | .compressed => { |
| 826 | 928 | if (src.len < block_size) return error.MalformedBlockSize; |
| 827 | 929 | var bytes_read: usize = 0; |
| 828 | | const literals = try decodeLiteralsSection(src, &bytes_read); |
| 829 | | const sequences_header = try decodeSequencesHeader(src[bytes_read..], &bytes_read); |
| 930 | const literals = decodeLiteralsSection(src, &bytes_read) catch return error.MalformedCompressedBlock; |
| 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 | 937 | var bytes_written: usize = 0; |
| 834 | 938 | if (sequences_header.sequence_count > 0) { |
| 835 | 939 | const bit_stream_bytes = src[bytes_read..block_size]; |
| 836 | 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 | 945 | var sequence_size_limit = block_size_max; |
| 842 | 946 | var i: usize = 0; |
| 843 | 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 | 949 | dest, |
| 846 | 950 | literals, |
| 847 | 951 | &bit_stream, |
| 848 | 952 | sequence_size_limit, |
| 849 | 953 | i == sequences_header.sequence_count - 1, |
| 850 | | ); |
| 954 | ) catch return error.MalformedCompressedBlock; |
| 851 | 955 | bytes_written += decompressed_size; |
| 852 | 956 | sequence_size_limit -= decompressed_size; |
| 853 | 957 | } |
| ... | ... | @@ -857,7 +961,8 @@ pub fn decodeBlockRingBuffer( |
| 857 | 961 | |
| 858 | 962 | if (decode_state.literal_written_count < literals.header.regenerated_size) { |
| 859 | 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 | 966 | bytes_written += len; |
| 862 | 967 | } |
| 863 | 968 | |
| ... | ... | @@ -866,7 +971,7 @@ pub fn decodeBlockRingBuffer( |
| 866 | 971 | consumed_count.* += bytes_read; |
| 867 | 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 | 1006 | } else return header.content_size; |
| 902 | 1007 | } |
| 903 | 1008 | |
| 1009 | const InvalidBit = error{ UnusedBitSet, ReservedBitSet }; |
| 904 | 1010 | /// Decode the header of a Zstandard frame. Returns `error.UnusedBitSet` or |
| 905 | 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 | 1013 | const descriptor = @bitCast(frame.ZStandard.Header.Descriptor, src[0]); |
| 908 | 1014 | |
| 909 | 1015 | if (descriptor.unused) return error.UnusedBitSet; |
| ... | ... | @@ -958,7 +1064,10 @@ pub fn decodeBlockHeader(src: *const [3]u8) frame.ZStandard.Block.Header { |
| 958 | 1064 | |
| 959 | 1065 | /// Decode a `LiteralsSection` from `src`, incrementing `consumed_count` by the |
| 960 | 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 | 1071 | var bytes_read: usize = 0; |
| 963 | 1072 | const header = try decodeLiteralsHeader(src, &bytes_read); |
| 964 | 1073 | switch (header.block_type) { |
| ... | ... | @@ -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 | 1151 | var bytes_read: usize = 0; |
| 1037 | 1152 | bytes_read += 1; |
| 1038 | 1153 | if (src.len == 0) return error.MalformedHuffmanTree; |
| ... | ... | @@ -1049,22 +1164,25 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H |
| 1049 | 1164 | var bit_reader = bitReader(counting_reader.reader()); |
| 1050 | 1165 | |
| 1051 | 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 | 1171 | const accuracy_log = std.math.log2_int_ceil(usize, table_size); |
| 1054 | 1172 | |
| 1055 | 1173 | const start_index = std.math.cast(usize, 1 + counting_reader.bytes_read) orelse return error.MalformedHuffmanTree; |
| 1056 | 1174 | var huff_data = src[start_index .. compressed_size + 1]; |
| 1057 | 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 | 1178 | var i: usize = 0; |
| 1061 | | var even_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log); |
| 1062 | | var odd_state: u32 = try huff_bits.readBitsNoEof(u32, accuracy_log); |
| 1179 | var even_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree; |
| 1180 | var odd_state: u32 = huff_bits.readBitsNoEof(u32, accuracy_log) catch return error.MalformedHuffmanTree; |
| 1063 | 1181 | |
| 1064 | 1182 | while (i < 255) { |
| 1065 | 1183 | const even_data = entries[even_state]; |
| 1066 | 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 | 1186 | weights[i] = std.math.cast(u4, even_data.symbol) orelse return error.MalformedHuffmanTree; |
| 1069 | 1187 | i += 1; |
| 1070 | 1188 | if (read_bits < even_data.bits) { |
| ... | ... | @@ -1076,7 +1194,7 @@ fn decodeHuffmanTree(src: []const u8, consumed_count: *usize) !LiteralsSection.H |
| 1076 | 1194 | |
| 1077 | 1195 | read_bits = 0; |
| 1078 | 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 | 1198 | weights[i] = std.math.cast(u4, odd_data.symbol) orelse return error.MalformedHuffmanTree; |
| 1081 | 1199 | i += 1; |
| 1082 | 1200 | if (read_bits < odd_data.bits) { |
| ... | ... | @@ -1177,8 +1295,8 @@ fn lessThanByWeight( |
| 1177 | 1295 | } |
| 1178 | 1296 | |
| 1179 | 1297 | /// Decode a literals section header. |
| 1180 | | pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSection.Header { |
| 1181 | | if (src.len == 0) return error.MalformedLiteralsSection; |
| 1298 | pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) error{MalformedLiteralsHeader}!LiteralsSection.Header { |
| 1299 | if (src.len == 0) return error.MalformedLiteralsHeader; |
| 1182 | 1300 | const byte0 = src[0]; |
| 1183 | 1301 | const block_type = @intToEnum(LiteralsSection.BlockType, byte0 & 0b11); |
| 1184 | 1302 | const size_format = @intCast(u2, (byte0 & 0b1100) >> 2); |
| ... | ... | @@ -1243,8 +1361,11 @@ pub fn decodeLiteralsHeader(src: []const u8, consumed_count: *usize) !LiteralsSe |
| 1243 | 1361 | } |
| 1244 | 1362 | |
| 1245 | 1363 | /// Decode a sequences section header. |
| 1246 | | pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !SequencesSection.Header { |
| 1247 | | if (src.len == 0) return error.MalformedSequencesSection; |
| 1364 | pub fn decodeSequencesHeader( |
| 1365 | src: []const u8, |
| 1366 | consumed_count: *usize, |
| 1367 | ) error{ MalformedSequencesHeader, ReservedBitSet }!SequencesSection.Header { |
| 1368 | if (src.len == 0) return error.MalformedSequencesHeader; |
| 1248 | 1369 | var sequence_count: u24 = undefined; |
| 1249 | 1370 | |
| 1250 | 1371 | var bytes_read: usize = 0; |
| ... | ... | @@ -1262,16 +1383,16 @@ pub fn decodeSequencesHeader(src: []const u8, consumed_count: *usize) !Sequences |
| 1262 | 1383 | sequence_count = byte0; |
| 1263 | 1384 | bytes_read += 1; |
| 1264 | 1385 | } else if (byte0 < 255) { |
| 1265 | | if (src.len < 2) return error.MalformedSequencesSection; |
| 1386 | if (src.len < 2) return error.MalformedSequencesHeader; |
| 1266 | 1387 | sequence_count = (@as(u24, (byte0 - 128)) << 8) + src[1]; |
| 1267 | 1388 | bytes_read += 2; |
| 1268 | 1389 | } else { |
| 1269 | | if (src.len < 3) return error.MalformedSequencesSection; |
| 1390 | if (src.len < 3) return error.MalformedSequencesHeader; |
| 1270 | 1391 | sequence_count = src[1] + (@as(u24, src[2]) << 8) + 0x7F00; |
| 1271 | 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 | 1396 | const compression_modes = src[bytes_read]; |
| 1276 | 1397 | bytes_read += 1; |
| 1277 | 1398 | |
| ... | ... | @@ -1441,17 +1562,17 @@ pub const ReverseBitReader = struct { |
| 1441 | 1562 | byte_reader: ReversedByteReader, |
| 1442 | 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 | 1566 | self.byte_reader = ReversedByteReader.init(bytes); |
| 1446 | 1567 | self.bit_reader = std.io.bitReader(.Big, self.byte_reader.reader()); |
| 1447 | 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 | 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 | 1576 | return try self.bit_reader.readBits(U, num_bits, out_bits); |
| 1456 | 1577 | } |
| 1457 | 1578 | |