from stdint import * import zlib.zchecksum as zchecksum import zlib.zdef as zdef import zlib.zhuff as zhuff import string import memhub import t, c @t.Object class zinflate_stream: pool: memhub.MemManager | t.CPtr reader: zdef.zbit_reader window: BYTEPTR window_pos: t.CInt window_size: t.CInt wbits: t.CInt is_finished: t.CInt is_initialized: t.CInt header_parsed: t.CInt is_gzip: t.CInt adler: t.CUInt32T crc: t.CUInt32T output: BYTEPTR output_len: t.CSizeT output_cap: t.CSizeT input_data: t.CConst | BYTEPTR input_len: t.CSizeT input_pos: t.CSizeT max_length: t.CSizeT def __init__(self): pass def output_byte(self, byte: BYTE): if self.max_length > 0 and self.output_len >= self.max_length: return if self.output_len >= self.output_cap: new_cap: t.CSizeT = self.output_cap * 2 if new_cap < 256: new_cap = 256 new_buf: BYTEPTR = BYTEPTR(zdef.zdef_alloc(self.pool, new_cap)) if not new_buf: return memcpy(new_buf, self.output, self.output_len) zdef.zdef_free(self.pool, self.output) self.output = new_buf self.output_cap = new_cap self.output[self.output_len] = byte self.output_len += 1 self.window[self.window_pos & (zdef.ZDEFLATE_WINDOW_SIZE - 1)] = byte self.window_pos += 1 def parse_header(self) -> t.CInt: if self.header_parsed: return 0 if self.wbits < 0: self.header_parsed = 1 self.is_gzip = 0 return 0 if self.input_pos + 2 > self.input_len: return -3 b0: BYTE = self.input_data[self.input_pos] b1: BYTE = self.input_data[self.input_pos + 1] if b0 == 0x1F and b1 == 0x8B: self.is_gzip = 1 if self.input_pos + 10 > self.input_len: return -3 self.input_pos += 10 flg: BYTE = self.input_data[self.input_pos - 6] if flg & 0x04: if self.input_pos + 2 > self.input_len: return -3 xlen: UINT = self.input_data[self.input_pos] | (self.input_data[self.input_pos + 1] << 8) self.input_pos += 2 + xlen if flg & 0x08: while self.input_pos < self.input_len and self.input_data[self.input_pos] != 0: self.input_pos += 1 self.input_pos += 1 if flg & 0x10: while self.input_pos < self.input_len and self.input_data[self.input_pos] != 0: self.input_pos += 1 self.input_pos += 1 if flg & 0x02: self.input_pos += 2 if self.input_pos > self.input_len: return -3 else: self.is_gzip = 0 if (b0 * 256 + b1) % 31 != 0: return -3 cm: t.CInt = b0 & 0x0F if cm != 8: return -3 self.input_pos += 2 self.header_parsed = 1 return 0 def read_bits_from_input(self, nbits: t.CInt, out: UINTPTR) -> t.CInt: val: UINT = 0 for i in range(nbits): if self.input_pos >= self.input_len: return -3 if self.input_data[self.input_pos] & (UINT(1) << (self.reader.bit_pos)): val |= (UINT(1) << i) self.reader.bit_pos += 1 if self.reader.bit_pos == 8: self.reader.bit_pos = 0 self.input_pos += 1 c.Set(c.Deref(out), val) return 0 def read_huffman(self, dt: zhuff.zhuff_decode_tree | t.CPtr) -> t.CInt: node: int = dt.root while dt.nodes[node].symbol == -1: bit: UINT if self.read_bits_from_input(1, c.Addr(bit)) != 0: return -1 dir: int = t.CInt(bit) if dt.nodes[node].children[dir] == -1: return -3 node = dt.nodes[node].children[dir] return dt.nodes[node].symbol def inflate_block_stored(self) -> t.CInt: if self.reader.bit_pos != 0: self.reader.bit_pos = 0 self.input_pos += 1 if self.input_pos + 4 > self.input_len: return -3 length: UINT = self.input_data[self.input_pos] | (self.input_data[self.input_pos + 1] << 8) nlen: UINT = self.input_data[self.input_pos + 2] | (self.input_data[self.input_pos + 3] << 8) self.input_pos += 4 if (length ^ nlen) != 0xFFFF: return -3 if self.input_pos + length > self.input_len: return -3 i: t.CUnsignedInt for i in range(length): self.output_byte(self.input_data[self.input_pos]) self.input_pos += 1 return 0 def inflate_block_fixed(self) -> t.CInt: lit_tree = zhuff.zhuff_tree() dist_tree = zhuff.zhuff_tree() lit_tree.build_fixed_lit_tree() dist_tree.build_fixed_dist_tree() lit_dt: zhuff.zhuff_decode_tree | t.CPtr = zdef.zdef_alloc(self.pool, zhuff.zhuff_decode_tree.__sizeof__()) dist_dt: zhuff.zhuff_decode_tree | t.CPtr = zdef.zdef_alloc(self.pool, zhuff.zhuff_decode_tree.__sizeof__()) if not lit_dt or not dist_dt: if lit_dt: zdef.zdef_free(self.pool, lit_dt) if dist_dt: zdef.zdef_free(self.pool, dist_dt) return -4 memset(lit_dt, 0, zhuff.zhuff_decode_tree.__sizeof__()) memset(dist_dt, 0, zhuff.zhuff_decode_tree.__sizeof__()) lit_dt.build_decode_tree(c.Addr(lit_tree)) dist_dt.build_decode_tree(c.Addr(dist_tree)) result: t.CInt = 0 while True: if self.max_length > 0 and self.output_len >= self.max_length: break sym: t.CInt = self.read_huffman(lit_dt) if sym < 0: result = -3 break if sym == 256: break if sym < 256: self.output_byte(BYTE(sym)) else: len_idx: t.CInt = sym - 257 if len_idx < 0 or len_idx >= 29: result = -3 break length: t.CInt = zdef.zdeflate_len_base[len_idx] extra: t.CInt = zdef.zdeflate_len_extra_bits[len_idx] if extra > 0: extra_val: UINT if self.read_bits_from_input(extra, c.Addr(extra_val)) != 0: result = -3 break length += extra_val dist_sym: t.CInt = self.read_huffman(dist_dt) if dist_sym < 0 or dist_sym >= 30: result = -3 break dist: t.CInt = zdef.zdeflate_dist_base[dist_sym] extra = zdef.zdeflate_dist_extra_bits[dist_sym] if extra > 0: extra_val: UINT if self.read_bits_from_input(extra, c.Addr(extra_val)) != 0: result = -3 break dist += extra_val for i in range(length): if self.max_length > 0 and self.output_len >= self.max_length: break src_pos: t.CInt = (self.window_pos - dist) & (zdef.ZDEFLATE_WINDOW_SIZE - 1) byte: BYTE = self.window[src_pos] self.output_byte(byte) zdef.zdef_free(self.pool, lit_dt) zdef.zdef_free(self.pool, dist_dt) return result def inflate_block_dynamic(self) -> t.CInt: hlit_val: UINT hdist_val: UINT hclen_val: UINT if self.read_bits_from_input(5, c.Addr(hlit_val)) != 0: return -3 if self.read_bits_from_input(5, c.Addr(hdist_val)) != 0: return -3 if self.read_bits_from_input(4, c.Addr(hclen_val)) != 0: return -3 hlit: t.CInt = t.CInt(hlit_val) + 257 hdist: t.CInt = t.CInt(hdist_val) + 1 hclen: t.CInt = t.CInt(hclen_val) + 4 cl_lengths: t.CArray[t.CInt, 19] memset(cl_lengths, 0, cl_lengths.__sizeof__()) for i in range(hclen): len_val: UINT if self.read_bits_from_input(3, c.Addr(len_val)) != 0: return -3 cl_lengths[zdef.zdeflate_codelen_order[i]] = t.CInt(len_val) cl_tree = zhuff.zhuff_tree() cl_tree.build_tree_from_lengths(cl_lengths, 19, 7) cl_dt = zhuff.zhuff_decode_tree() cl_dt.build_decode_tree(c.Addr(cl_tree)) total: t.CInt = hlit + hdist all_lengths: t.CInt | t.CPtr = zdef.zdef_alloc(self.pool, total * int.__sizeof__()) if not all_lengths: return -4 idx: t.CInt = 0 while idx < total: sym: t.CInt = self.read_huffman(c.Addr(cl_dt)) if sym < 0 or sym > 18: zdef.zdef_free(self.pool, all_lengths) return -3 if sym < 16: all_lengths[idx] = sym idx += 1 elif sym == 16: rep: UINT if self.read_bits_from_input(2, c.Addr(rep)) != 0: zdef.zdef_free(self.pool, all_lengths) return -3 rep += 3 if idx == 0 or idx + rep > total: zdef.zdef_free(self.pool, all_lengths) return -3 prev: t.CInt = all_lengths[idx - 1] i: t.CUnsignedInt for i in range(rep): all_lengths[idx] = prev idx += 1 elif sym == 17: rep: t.CUnsignedInt if self.read_bits_from_input(3, c.Addr(rep)) != 0: zdef.zdef_free(self.pool, all_lengths) return -3 rep += 3 if idx + rep > total: zdef.zdef_free(self.pool, all_lengths) return -3 i: t.CUnsignedInt for i in range(rep): all_lengths[idx] = 0 idx += 1 else: rep: t.CUnsignedInt if self.read_bits_from_input(7, c.Addr(rep)) != 0: zdef.zdef_free(self.pool, all_lengths) return -3 rep += 11 if idx + rep > total: zdef.zdef_free(self.pool, all_lengths) return -3 i: t.CUnsignedInt for i in range(rep): all_lengths[idx] = 0 idx += 1 lit_tree = zhuff.zhuff_tree() dist_tree = zhuff.zhuff_tree() lit_lengths: t.CArray[t.CInt, 288] memset(lit_lengths, 0, lit_lengths.__sizeof__()) for i in range(hlit): lit_lengths[i] = all_lengths[i] lit_tree.build_tree_from_lengths(lit_lengths, hlit, zdef.ZDEFLATE_MAX_BITS) dist_lengths: t.CArray[t.CInt, 32] memset(dist_lengths, 0, dist_lengths.__sizeof__()) for i in range(hdist): dist_lengths[i] = all_lengths[hlit + i] dist_tree.build_tree_from_lengths(dist_lengths, hdist, zdef.ZDEFLATE_MAX_BITS) zdef.zdef_free(self.pool, all_lengths) lit_dt = zhuff.zhuff_decode_tree() dist_dt = zhuff.zhuff_decode_tree() lit_dt.build_decode_tree(c.Addr(lit_tree)) dist_dt.build_decode_tree(c.Addr(dist_tree)) while True: if self.max_length > 0 and self.output_len >= self.max_length: return 0 sym: t.CInt = self.read_huffman(c.Addr(lit_dt)) if sym < 0: return -3 if sym == 256: return 0 if sym < 256: self.output_byte(BYTE(sym)) else: len_idx: t.CInt = sym - 257 if len_idx < 0 or len_idx >= 29: return -3 length: t.CInt = zdef.zdeflate_len_base[len_idx] extra: t.CInt = zdef.zdeflate_len_extra_bits[len_idx] if extra > 0: extra_val: UINT if self.read_bits_from_input(extra, c.Addr(extra_val)) != 0: return -3 length += extra_val dist_sym: t.CInt = self.read_huffman(c.Addr(dist_dt)) if dist_sym < 0 or dist_sym >= 30: return -3 dist: t.CInt = zdef.zdeflate_dist_base[dist_sym] extra = zdef.zdeflate_dist_extra_bits[dist_sym] if extra > 0: extra_val: UINT if self.read_bits_from_input(extra, c.Addr(extra_val)) != 0: return -3 dist += extra_val for i in range(length): if self.max_length > 0 and self.output_len >= self.max_length: return 0 src_pos: t.CInt = (self.window_pos - dist) & (zdef.ZDEFLATE_WINDOW_SIZE - 1) byte: BYTE = self.window[src_pos] self.output_byte(byte) def inflate_stream(self) -> t.CInt: err: t.CInt = self.parse_header() if err != 0: return err while not self.is_finished: bfinal_val: UINT btype_val: UINT if self.read_bits_from_input(1, c.Addr(bfinal_val)) != 0: return -3 if self.read_bits_from_input(2, c.Addr(btype_val)) != 0: return -3 bfinal: t.CInt = t.CInt(bfinal_val) btype: t.CInt = t.CInt(btype_val) if btype == 0: err = self.inflate_block_stored() elif btype == 1: err = self.inflate_block_fixed() elif btype == 2: err = self.inflate_block_dynamic() else: return -3 if err != 0: return err if bfinal: self.is_finished = 1 if self.max_length > 0 and self.output_len >= self.max_length: return 0 if self.is_gzip: if self.reader.bit_pos != 0: self.reader.bit_pos = 0 self.input_pos += 1 if self.input_pos + 8 <= self.input_len: self.crc = ((t.CUInt32T(self.input_data[self.input_pos + 0]) << 0) | (t.CUInt32T(self.input_data[self.input_pos + 1]) << 8) | (t.CUInt32T(self.input_data[self.input_pos + 2]) << 16) | (t.CUInt32T(self.input_data[self.input_pos + 3]) << 24)) self.input_pos += 8 elif self.wbits >= 0: if self.reader.bit_pos != 0: self.reader.bit_pos = 0 self.input_pos += 1 if self.input_pos + 4 <= self.input_len: self.adler = ((t.CUInt32T(self.input_data[self.input_pos + 0]) << 24) | (t.CUInt32T(self.input_data[self.input_pos + 1]) << 16) | (t.CUInt32T(self.input_data[self.input_pos + 2]) << 8) | (t.CUInt32T(self.input_data[self.input_pos + 3]) << 0)) self.input_pos += 4 return 0 def decompress(self, data: UINT8PTR, length: t.CSizeT, max_length: t.CSizeT, out: t.CUInt8T | t.CPtr[t.CPtr], out_len: t.CSizeT | t.CPtr) -> t.CInt: if not self.is_initialized: return -2 self.input_data = data self.input_len = length self.input_pos = 0 self.reader.bit_pos = 0 self.output_len = 0 self.max_length = max_length err: t.CInt = self.inflate_stream() if err != 0: return err c.Set(c.Deref(out_len), self.output_len) if self.output_len == 0: c.Set(c.Deref(out), UINT8PTR(zdef.zdef_alloc(self.pool, 1))) else: dst: UINT8PTR = UINT8PTR(zdef.zdef_alloc(self.pool, self.output_len)) c.Set(c.Deref(out), dst) if dst: memcpy(dst, self.output, self.output_len) return 0 def set_dictionary(self, dict: UINT8PTR, length: t.CSizeT) -> t.CInt: if not dict: return -2 use: t.CSizeT = length if use > zdef.ZDEFLATE_WINDOW_SIZE: dict += length - zdef.ZDEFLATE_WINDOW_SIZE use = zdef.ZDEFLATE_WINDOW_SIZE memcpy(self.window, dict, use) self.window_pos = t.CInt(use) return 0 def copy(self) -> zinflate_stream | t.CPtr: z: zinflate_stream | t.CPtr = zdef.zdef_alloc(self.pool, zinflate_stream.__sizeof__()) if not z: return None memcpy(z, self, zinflate_stream.__sizeof__()) z.window = BYTEPTR(zdef.zdef_alloc(self.pool, zdef.ZDEFLATE_WINDOW_SIZE)) z.output = BYTEPTR(zdef.zdef_alloc(self.pool, self.output_cap)) if not z.window or not z.output: zinflate_stream.destroy(z) return None memcpy(z.window, self.window, zdef.ZDEFLATE_WINDOW_SIZE) memcpy(z.output, self.output, self.output_cap) return z def destroy(self): if not self: return if self.window: zdef.zdef_free(self.pool, self.window) if self.output: zdef.zdef_free(self.pool, self.output) zdef.zdef_free(self.pool, self) def zinflate_one_shot(pool: memhub.MemManager | t.CPtr, data: UINT8PTR, length: t.CSizeT, wbits: t.CInt, bufsize: t.CSizeT, out_len: t.CSizeT | t.CPtr) -> t.CUInt8T | t.CPtr: s: zinflate_stream | t.CPtr = zinflate_create(pool, wbits) if not s: return None out: t.CUInt8T | t.CPtr = None err: t.CInt = s.decompress(data, length, 0, c.Addr(out), out_len) s.destroy() if err != 0: return None return out def zinflate_create(pool: memhub.MemManager | t.CPtr, wbits: t.CInt) -> zinflate_stream | t.CPtr: s: zinflate_stream | t.CPtr = zdef.zdef_alloc(pool, zinflate_stream.__sizeof__()) if not s: return None memset(s, 0, zinflate_stream.__sizeof__()) s.pool = pool s.wbits = wbits s.is_initialized = 1 s.adler = 1 s.window = BYTEPTR(zdef.zdef_alloc(pool, zdef.ZDEFLATE_WINDOW_SIZE)) s.output = BYTEPTR(zdef.zdef_alloc(pool, 256)) s.output_cap = 256 s.output_len = 0 s.max_length = 0 if not s.window or not s.output: s.destroy() return None return s