| ... | @@ -148,8 +148,8 @@ pub const DecodeState = struct { | ... | @@ -148,8 +148,8 @@ pub const DecodeState = struct { |
| 148 | } | 148 | } |
| 149 | | 149 | |
| 150 | fn updateRepeatOffset(self: *DecodeState, offset: u32) void { | 150 | fn updateRepeatOffset(self: *DecodeState, offset: u32) void { |
| 151 | std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[1]); | 151 | self.repeat_offsets[2] = self.repeat_offsets[1]; |
| 152 | std.mem.swap(u32, &self.repeat_offsets[0], &self.repeat_offsets[2]); | 152 | self.repeat_offsets[1] = self.repeat_offsets[0]; |
| 153 | self.repeat_offsets[0] = offset; | 153 | self.repeat_offsets[0] = offset; |
| 154 | } | 154 | } |
| 155 | | 155 | |
| ... | @@ -238,18 +238,22 @@ pub const DecodeState = struct { | ... | @@ -238,18 +238,22 @@ pub const DecodeState = struct { |
| 238 | fn nextSequence( | 238 | fn nextSequence( |
| 239 | self: *DecodeState, | 239 | self: *DecodeState, |
| 240 | bit_reader: *readers.ReverseBitReader, | 240 | bit_reader: *readers.ReverseBitReader, |
| 241 | ) error{ OffsetCodeTooLarge, EndOfStream }!Sequence { | 241 | ) error{ InvalidBitStream, EndOfStream }!Sequence { |
| 242 | const raw_code = self.getCode(.offset); | 242 | const raw_code = self.getCode(.offset); |
| 243 | const offset_code = std.math.cast(u5, raw_code) orelse { | 243 | const offset_code = std.math.cast(u5, raw_code) orelse { |
| 244 | return error.OffsetCodeTooLarge; | 244 | return error.InvalidBitStream; |
| 245 | }; | 245 | }; |
| 246 | const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code); | 246 | const offset_value = (@as(u32, 1) << offset_code) + try bit_reader.readBitsNoEof(u32, offset_code); |
| 247 | | 247 | |
| 248 | const match_code = self.getCode(.match); | 248 | const match_code = self.getCode(.match); |
| | 249 | if (match_code >= types.compressed_block.match_length_code_table.len) |
| | 250 | return error.InvalidBitStream; |
| 249 | const match = types.compressed_block.match_length_code_table[match_code]; | 251 | const match = types.compressed_block.match_length_code_table[match_code]; |
| 250 | const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]); | 252 | const match_length = match[0] + try bit_reader.readBitsNoEof(u32, match[1]); |
| 251 | | 253 | |
| 252 | const literal_code = self.getCode(.literal); | 254 | const literal_code = self.getCode(.literal); |
| | 255 | if (literal_code >= types.compressed_block.literals_length_code_table.len) |
| | 256 | return error.InvalidBitStream; |
| 253 | const literal = types.compressed_block.literals_length_code_table[literal_code]; | 257 | const literal = types.compressed_block.literals_length_code_table[literal_code]; |
| 254 | const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]); | 258 | const literal_length = literal[0] + try bit_reader.readBitsNoEof(u32, literal[1]); |
| 255 | | 259 | |
| ... | @@ -269,6 +273,8 @@ pub const DecodeState = struct { | ... | @@ -269,6 +273,8 @@ pub const DecodeState = struct { |
| 269 | break :offset self.useRepeatOffset(offset_value - 1); | 273 | break :offset self.useRepeatOffset(offset_value - 1); |
| 270 | }; | 274 | }; |
| 271 | | 275 | |
| | 276 | if (offset == 0) return error.InvalidBitStream; |
| | 277 | |
| 272 | return .{ | 278 | return .{ |
| 273 | .literal_length = literal_length, | 279 | .literal_length = literal_length, |
| 274 | .match_length = match_length, | 280 | .match_length = match_length, |
| ... | @@ -308,7 +314,7 @@ pub const DecodeState = struct { | ... | @@ -308,7 +314,7 @@ pub const DecodeState = struct { |
| 308 | } | 314 | } |
| 309 | | 315 | |
| 310 | const DecodeSequenceError = error{ | 316 | const DecodeSequenceError = error{ |
| 311 | OffsetCodeTooLarge, | 317 | InvalidBitStream, |
| 312 | EndOfStream, | 318 | EndOfStream, |
| 313 | MalformedSequence, | 319 | MalformedSequence, |
| 314 | MalformedFseBits, | 320 | MalformedFseBits, |
| ... | @@ -326,7 +332,7 @@ pub const DecodeState = struct { | ... | @@ -326,7 +332,7 @@ pub const DecodeState = struct { |
| 326 | /// - `error.UnexpectedEndOfLiteralStream` if the decoder state's literal | 332 | /// - `error.UnexpectedEndOfLiteralStream` if the decoder state's literal |
| 327 | /// streams do not contain enough literals for the sequence (this may | 333 | /// streams do not contain enough literals for the sequence (this may |
| 328 | /// mean the literal stream or the sequence is malformed). | 334 | /// mean the literal stream or the sequence is malformed). |
| 329 | /// - `error.OffsetCodeTooLarge` if an invalid offset code is found | 335 | /// - `error.InvalidBitStream` if the FSE sequence bitstream is malformed |
| 330 | /// - `error.EndOfStream` if `bit_reader` does not contain enough bits | 336 | /// - `error.EndOfStream` if `bit_reader` does not contain enough bits |
| 331 | pub fn decodeSequenceSlice( | 337 | pub fn decodeSequenceSlice( |
| 332 | self: *DecodeState, | 338 | self: *DecodeState, |
| ... | @@ -608,9 +614,9 @@ pub fn decodeBlock( | ... | @@ -608,9 +614,9 @@ pub fn decodeBlock( |
| 608 | .compressed => { | 614 | .compressed => { |
| 609 | if (src.len < block_size) return error.MalformedBlockSize; | 615 | if (src.len < block_size) return error.MalformedBlockSize; |
| 610 | var bytes_read: usize = 0; | 616 | var bytes_read: usize = 0; |
| 611 | const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch | 617 | const literals = decodeLiteralsSectionSlice(src[0..block_size], &bytes_read) catch |
| 612 | return error.MalformedCompressedBlock; | 618 | return error.MalformedCompressedBlock; |
| 613 | var fbs = std.io.fixedBufferStream(src[bytes_read..]); | 619 | var fbs = std.io.fixedBufferStream(src[bytes_read..block_size]); |
| 614 | const fbs_reader = fbs.reader(); | 620 | const fbs_reader = fbs.reader(); |
| 615 | const sequences_header = decodeSequencesHeader(fbs_reader) catch | 621 | const sequences_header = decodeSequencesHeader(fbs_reader) catch |
| 616 | return error.MalformedCompressedBlock; | 622 | return error.MalformedCompressedBlock; |
| ... | @@ -695,9 +701,9 @@ pub fn decodeBlockRingBuffer( | ... | @@ -695,9 +701,9 @@ pub fn decodeBlockRingBuffer( |
| 695 | .compressed => { | 701 | .compressed => { |
| 696 | if (src.len < block_size) return error.MalformedBlockSize; | 702 | if (src.len < block_size) return error.MalformedBlockSize; |
| 697 | var bytes_read: usize = 0; | 703 | var bytes_read: usize = 0; |
| 698 | const literals = decodeLiteralsSectionSlice(src, &bytes_read) catch | 704 | const literals = decodeLiteralsSectionSlice(src[0..block_size], &bytes_read) catch |
| 699 | return error.MalformedCompressedBlock; | 705 | return error.MalformedCompressedBlock; |
| 700 | var fbs = std.io.fixedBufferStream(src[bytes_read..]); | 706 | var fbs = std.io.fixedBufferStream(src[bytes_read..block_size]); |
| 701 | const fbs_reader = fbs.reader(); | 707 | const fbs_reader = fbs.reader(); |
| 702 | const sequences_header = decodeSequencesHeader(fbs_reader) catch | 708 | const sequences_header = decodeSequencesHeader(fbs_reader) catch |
| 703 | return error.MalformedCompressedBlock; | 709 | return error.MalformedCompressedBlock; |
| ... | @@ -894,7 +900,8 @@ pub fn decodeLiteralsSectionSlice( | ... | @@ -894,7 +900,8 @@ pub fn decodeLiteralsSectionSlice( |
| 894 | else | 900 | else |
| 895 | null; | 901 | null; |
| 896 | const huffman_tree_size = bytes_read - huffman_tree_start; | 902 | const huffman_tree_size = bytes_read - huffman_tree_start; |
| 897 | const total_streams_size = @as(usize, header.compressed_size.?) - huffman_tree_size; | 903 | const total_streams_size = std.math.sub(usize, header.compressed_size.?, huffman_tree_size) catch |
| | 904 | return error.MalformedLiteralsSection; |
| 898 | | 905 | |
| 899 | if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection; | 906 | if (src.len < bytes_read + total_streams_size) return error.MalformedLiteralsSection; |
| 900 | const stream_data = src[bytes_read .. bytes_read + total_streams_size]; | 907 | const stream_data = src[bytes_read .. bytes_read + total_streams_size]; |
| ... | @@ -940,8 +947,9 @@ pub fn decodeLiteralsSection( | ... | @@ -940,8 +947,9 @@ pub fn decodeLiteralsSection( |
| 940 | try huffman.decodeHuffmanTree(counting_reader.reader(), buffer) | 947 | try huffman.decodeHuffmanTree(counting_reader.reader(), buffer) |
| 941 | else | 948 | else |
| 942 | null; | 949 | null; |
| 943 | const huffman_tree_size = counting_reader.bytes_read; | 950 | const huffman_tree_size = @intCast(usize, counting_reader.bytes_read); |
| 944 | const total_streams_size = @as(usize, header.compressed_size.?) - @intCast(usize, huffman_tree_size); | 951 | const total_streams_size = std.math.sub(usize, header.compressed_size.?, huffman_tree_size) catch |
| | 952 | return error.MalformedLiteralsSection; |
| 945 | | 953 | |
| 946 | if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall; | 954 | if (total_streams_size > buffer.len) return error.LiteralsBufferTooSmall; |
| 947 | try source.readNoEof(buffer[0..total_streams_size]); | 955 | try source.readNoEof(buffer[0..total_streams_size]); |