from stdint import * import zlib.zchecksum as zchecksum import zlib.zhuff as zhuff import zlib.zdef as zdef import string import memhub import t, c @t.Object class zdeflate_stream: pool: memhub.MemManager | t.CPtr writer: zdef.zbit_writer window: BYTE | t.CPtr window_pos: t.CInt window_size: t.CInt hash_head: t.CInt | t.CPtr hash_prev: t.CInt | t.CPtr level: t.CInt strategy: t.CInt wbits: t.CInt is_finished: t.CInt adler: t.CUInt32T def find_match(self, data: BYTE | t.CPtr, data_len: t.CSizeT, pos: t.CSizeT, best_dist: t.CInt | t.CPtr) -> t.CInt: if pos + 2 >= data_len: return 0 h: t.CUnsignedInt = zdeflate_hash(data + pos) chain: t.CInt = self.hash_head[h] best_len: t.CInt = zdef.ZDEFLATE_MIN_MATCH - 1 c.Set(c.Deref(best_dist), 0) max_chain: t.CInt = 128 if (self.level >= 5) else (32 if (self.level >= 2) else 4) limit: t.CInt = (1 << self.wbits) if (self.wbits > 0) else zdef.ZDEFLATE_WINDOW_SIZE if limit > zdef.ZDEFLATE_WINDOW_SIZE: limit = zdef.ZDEFLATE_WINDOW_SIZE attempts: t.CInt = 0 while chain >= 0 and attempts < max_chain: dist: t.CInt = t.CInt(pos) - chain if dist <= 0 or dist > limit: break match_len: t.CInt = 0 max_len: t.CInt = t.CInt(data_len - pos) if max_len > zdef.ZDEFLATE_MAX_MATCH: max_len = zdef.ZDEFLATE_MAX_MATCH while match_len < max_len and data[pos + match_len] == data[chain + match_len]: match_len += 1 if match_len > best_len: best_len = match_len c.Set(c.Deref(best_dist), dist) if best_len >= zdef.ZDEFLATE_MAX_MATCH: break chain = self.hash_prev[chain & zdef.ZDEFLATE_WINDOW_MASK] if chain <= t.CInt(pos) - limit: break attempts += 1 return best_len if (best_len >= zdef.ZDEFLATE_MIN_MATCH) else 0 def update_hash(self, data: BYTE | t.CPtr, pos: t.CSizeT): if pos + 2 < pos: return h: t.CUnsignedInt = zdeflate_hash(data + pos) self.hash_prev[pos & zdef.ZDEFLATE_WINDOW_MASK] = self.hash_head[h] self.hash_head[h] = t.CInt(pos) def write_fixed_block(self, data: UINT8PTR, length: t.CSizeT, final: t.CInt): lit_tree = zhuff.zhuff_tree() dist_tree = zhuff.zhuff_tree() lit_tree.build_fixed_lit_tree() dist_tree.build_fixed_dist_tree() self.writer.write_bits(1 if final else 0, 1) self.writer.write_bits(1, 2) pos: t.CSizeT = 0 while pos < length: best_dist: t.CInt = 0 match_len: t.CInt = 0 if self.level > 0 and pos + 2 < length: match_len = self.find_match(data, length, pos, c.Addr(best_dist)) if match_len >= zdef.ZDEFLATE_MIN_MATCH: len_sym: t.CInt = zdeflate_len_to_symbol(match_len) lit_tree.encode_symbol(len_sym, c.Addr(self.writer)) extra: t.CInt = zdef.zdeflate_len_extra_bits[len_sym - 257] if extra > 0: self.writer.write_bits(match_len - zdef.zdeflate_len_base[len_sym - 257], extra) dist_sym: t.CInt = zdeflate_dist_to_symbol(best_dist) dist_tree.encode_symbol(dist_sym, c.Addr(self.writer)) extra = zdef.zdeflate_dist_extra_bits[dist_sym] if extra > 0: self.writer.write_bits(best_dist - zdef.zdeflate_dist_base[dist_sym], extra) for i in range(match_len): self.update_hash(data, pos + i) pos += match_len else: lit_tree.encode_symbol(data[pos], c.Addr(self.writer)) self.update_hash(data, pos) pos += 1 lit_tree.encode_symbol(256, c.Addr(self.writer)) def encode_block_data(self, data: UINT8PTR, length: t.CSizeT, lit_tree: zhuff.zhuff_tree | t.CPtr, dist_tree: zhuff.zhuff_tree | t.CPtr): pos: t.CSizeT = 0 while pos < length: best_dist: t.CInt = 0 match_len: t.CInt = 0 if self.level > 0 and pos + 2 < length: match_len = self.find_match(data, length, pos, c.Addr(best_dist)) if match_len >= zdef.ZDEFLATE_MIN_MATCH: len_sym: t.CInt = zdeflate_len_to_symbol(match_len) lit_tree.encode_symbol(len_sym, c.Addr(self.writer)) extra: t.CInt = zdef.zdeflate_len_extra_bits[len_sym - 257] if extra > 0: self.writer.write_bits(match_len - zdef.zdeflate_len_base[len_sym - 257], extra) dist_sym: t.CInt = zdeflate_dist_to_symbol(best_dist) dist_tree.encode_symbol(dist_sym, c.Addr(self.writer)) extra = zdef.zdeflate_dist_extra_bits[dist_sym] if extra > 0: self.writer.write_bits(best_dist - zdef.zdeflate_dist_base[dist_sym], extra) for j in range(match_len): self.update_hash(data, pos + j) pos += match_len else: lit_tree.encode_symbol(data[pos], c.Addr(self.writer)) self.update_hash(data, pos) pos += 1 lit_tree.encode_symbol(256, c.Addr(self.writer)) def write_dynamic_block(self, data: UINT8PTR, length: t.CSizeT, final: t.CInt): lit_freqs: t.CArray[t.CInt, 288] dist_freqs: t.CArray[t.CInt, 32] zdeflate_count_freqs(lit_freqs, dist_freqs, self, data, length) hlit: t.CInt = 286 while hlit > 257 and lit_freqs[hlit - 1] == 0: hlit -= 1 hdist: t.CInt = 30 while hdist > 1 and dist_freqs[hdist - 1] == 0: hdist -= 1 lit_tree = zhuff.zhuff_tree() dist_tree = zhuff.zhuff_tree() lit_tree.build_codes(lit_freqs, hlit, zdef.ZDEFLATE_MAX_BITS) dist_tree.build_codes(dist_freqs, hdist, zdef.ZDEFLATE_MAX_BITS) all_lengths: t.CArray[t.CInt, 288 + 32] for i in range(hlit): all_lengths[i] = lit_tree.codes[i].bits for i in range(hdist): all_lengths[hlit + i] = dist_tree.codes[i].bits total_lengths: t.CInt = hlit + hdist cl_freqs: t.CArray[t.CInt, 19] zdeflate_count_cl_freqs(all_lengths, total_lengths, cl_freqs) cl_tree = zhuff.zhuff_tree() cl_tree.build_codes(cl_freqs, 19, 7) hclen: int = 19 while hclen > 4 and cl_tree.codes[zdef.zdeflate_codelen_order[hclen - 1]].bits == 0: hclen -= 1 self.writer.write_bits(1 if final else 0, 1) self.writer.write_bits(2, 2) self.writer.write_bits(hlit - 257, 5) self.writer.write_bits(hdist - 1, 5) self.writer.write_bits(hclen - 4, 4) for j in range(hclen): self.writer.write_bits(cl_tree.codes[zdef.zdeflate_codelen_order[j]].bits, 3) zdeflate_write_cl_encoded(all_lengths, total_lengths, c.Addr(self.writer), c.Addr(cl_tree)) self.encode_block_data(data, length, c.Addr(lit_tree), c.Addr(dist_tree)) def compress_block(self, data: UINT8PTR, length: t.CSizeT, final: t.CInt): if self.level == 0: self.write_stored_block(data, length, final) elif self.strategy == 4 or length < 128: self.write_fixed_block(data, length, final) else: self.write_dynamic_block(data, length, final) def write_stored_block(self, data: UINT8PTR, length: t.CSizeT, final: t.CInt): self.writer.stored_block(data, length, final) def compress(self, data: UINT8PTR, length: t.CSizeT, out: t.CUInt8T | t.CPtr[t.CPtr], out_len: t.CSizeT | t.CPtr) -> t.CInt: if self.is_finished: return -2 if not data or length == 0: c.Set(c.Deref(out), None) c.Set(c.Deref(out_len), 0) return 0 self.adler = zchecksum.zchecksum_adler32(data, length, self.adler) memset(self.hash_head, -1, zdef.ZDEFLATE_HASH_SIZE * int.__sizeof__()) memset(self.hash_prev, -1, zdef.ZDEFLATE_WINDOW_SIZE * int.__sizeof__()) self.compress_block(data, length, 0) c.Set(c.Deref(out_len), self.writer.total()) c.Set(c.Deref(out), UINT8PTR(zdef.zdef_alloc(self.pool, c.Deref(out_len)))) if not c.Deref(out): return -4 memcpy(c.Deref(out), self.writer.buf, c.Deref(out_len)) return 0 def flush(self, mode: t.CInt, out: t.CUInt8T | t.CPtr[t.CPtr], out_len: t.CSizeT | t.CPtr) -> t.CInt: if mode == 4: if not self.is_finished: if self.level == 0: self.writer.write_bits(1, 1) self.writer.write_bits(0, 2) self.writer.align() else: lit_tree = zhuff.zhuff_tree() lit_tree.build_fixed_lit_tree() self.writer.write_bits(1, 1) self.writer.write_bits(1, 2) lit_tree.encode_symbol(256, c.Addr(self.writer)) self.is_finished = 1 if self.wbits >= 0: self.writer.align() adler: t.CUInt32T = self.adler self.writer.buf[self.writer.byte_pos + 0] = (adler >> 24) & 0xFF self.writer.buf[self.writer.byte_pos + 1] = (adler >> 16) & 0xFF self.writer.buf[self.writer.byte_pos + 2] = (adler >> 8) & 0xFF self.writer.buf[self.writer.byte_pos + 3] = (adler >> 0) & 0xFF self.writer.byte_pos += 4 elif mode == 2 or mode == 3: self.writer.write_bits(0, 1) self.writer.write_bits(0, 2) self.writer.align() c.Set(c.Deref(out_len), self.writer.total()) c.Set(c.Deref(out), UINT8PTR(zdef.zdef_alloc(self.pool, c.Deref(out_len)))) if not c.Deref(out): return -4 memcpy(c.Deref(out), self.writer.buf, c.Deref(out_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_size = t.CInt(use) self.window_pos = t.CInt(use) self.adler = zchecksum.zchecksum_adler32(dict, length, self.adler) return 0 def copy(self) -> zdeflate_stream | t.CPtr: z: zdeflate_stream | t.CPtr = zdef.zdef_alloc(self.pool, zdeflate_stream.__sizeof__()) if not z: return None memcpy(z, self, zdeflate_stream.__sizeof__()) z.writer.buf = BYTEPTR(zdef.zdef_alloc(self.pool, self.writer.cap)) if not z.writer.buf: zdef.zdef_free(self.pool, z) return None memcpy(z.writer.buf, self.writer.buf, self.writer.cap) z.window = BYTEPTR(zdef.zdef_alloc(self.pool, zdef.ZDEFLATE_WINDOW_SIZE)) z.hash_head = INTPTR(zdef.zdef_alloc(self.pool, zdef.ZDEFLATE_HASH_SIZE * int.__sizeof__())) z.hash_prev = INTPTR(zdef.zdef_alloc(self.pool, zdef.ZDEFLATE_WINDOW_SIZE * int.__sizeof__())) if not z.window or not z.hash_head or not z.hash_prev: if z.writer.buf: zdef.zdef_free(self.pool, z.writer.buf) zdef.zdef_free(self.pool, z) return None memcpy(z.window, self.window, zdef.ZDEFLATE_WINDOW_SIZE) memcpy(z.hash_head, self.hash_head, zdef.ZDEFLATE_HASH_SIZE * int.__sizeof__()) memcpy(z.hash_prev, self.hash_prev, zdef.ZDEFLATE_WINDOW_SIZE * int.__sizeof__()) return z def destroy(self): if self.writer.buf: zdef.zdef_free(self.pool, self.writer.buf) if self.window: zdef.zdef_free(self.pool, self.window) if self.hash_head: zdef.zdef_free(self.pool, self.hash_head) if self.hash_prev: zdef.zdef_free(self.pool, self.hash_prev) zdef.zdef_free(self.pool, self) def zdeflate_count_freqs(lit_freqs: INTPTR, dist_freqs: INTPTR, s: zdeflate_stream | t.CPtr, data: UINT8PTR, length: t.CSizeT): memset(lit_freqs, 0, 288 * int.__sizeof__()) memset(dist_freqs, 0, 32 * int.__sizeof__()) pos: t.CSizeT = 0 while pos < length: best_dist: t.CInt = 0 match_len: t.CInt = 0 if s.level > 0 and pos + 2 < length: match_len = zdeflate_stream.find_match(s, data, length, pos, c.Addr(best_dist)) if match_len >= zdef.ZDEFLATE_MIN_MATCH: len_sym: t.CInt = zdeflate_len_to_symbol(match_len) lit_freqs[len_sym] += 1 dist_freqs[zdeflate_dist_to_symbol(best_dist)] += 1 for i in range(match_len): zdeflate_stream.update_hash(s, data, pos + i) pos += match_len else: lit_freqs[data[pos]] += 1 zdeflate_stream.update_hash(s, data, pos) pos += 1 lit_freqs[256] = 1 def zdeflate_hash(p: BYTE | t.CPtr) -> t.CUnsignedInt: return (t.CUnsignedInt(p[0]) ^ (t.CUnsignedInt(p[1]) << 5) ^ (t.CUnsignedInt(p[2]) << 10)) & zdef.ZDEFLATE_HASH_MASK def zdeflate_len_to_symbol(length: t.CInt) -> t.CInt: for i in range(29): if length <= zdef.zdeflate_len_base[i] + (1 << zdef.zdeflate_len_extra_bits[i]) - 1: return 257 + i return 285 def zdeflate_dist_to_symbol(dist: t.CInt) -> t.CInt: for i in range(30): if dist <= zdef.zdeflate_dist_base[i] + (1 << zdef.zdeflate_dist_extra_bits[i]) - 1: return i return 29 def zdeflate_count_cl_freqs(all_lengths: INTPTR, total: t.CInt, cl_freqs: INTPTR): memset(cl_freqs, 0, 19 * int.__sizeof__()) i: t.CInt = 0 while i < total: if all_lengths[i] == 0: run: t.CInt = 0 while i + run < total and all_lengths[i + run] == 0: run += 1 while run > 0: if run >= 11: count: t.CInt = run if count > 138: count = 138 cl_freqs[18] += 1 run -= count elif run >= 3: count: t.CInt = run if count > 10: count = 10 cl_freqs[17] += 1 run -= count else: cl_freqs[0] += 1 run -= 1 while i < total and all_lengths[i] == 0: i += 1 else: cl_freqs[all_lengths[i]] += 1 val: t.CInt = all_lengths[i] i += 1 run: t.CInt = 0 while i + run < total and all_lengths[i + run] == val: run += 1 while run >= 3: count: t.CInt = run if count > 6: count = 6 cl_freqs[16] += 1 run -= count i += run def zdeflate_write_cl_encoded(all_lengths: INTPTR, total: t.CInt, w: zdef.zbit_writer | t.CPtr, cl_tree: zhuff.zhuff_tree | t.CPtr): i: t.CInt = 0 while i < total: if all_lengths[i] == 0: run: t.CInt = 0 while i + run < total and all_lengths[i + run] == 0: run += 1 while run > 0: if run >= 11: count: t.CInt = run if count > 138: count = 138 cl_tree.encode_symbol(18, w) w.write_bits(count - 11, 7) run -= count elif run >= 3: count: t.CInt = run if count > 10: count = 10 cl_tree.encode_symbol(17, w) w.write_bits(count - 3, 3) run -= count else: cl_tree.encode_symbol(0, w) run -= 1 while i < total and all_lengths[i] == 0: i += 1 else: cl_tree.encode_symbol(all_lengths[i], w) val: t.CInt = all_lengths[i] i += 1 run: t.CInt = 0 while i + run < total and all_lengths[i + run] == val: run += 1 while run >= 3: count: t.CInt = run if count > 6: count = 6 cl_tree.encode_symbol(16, w) w.write_bits(count - 3, 2) run -= count i += run def zdeflate_create(pool: memhub.MemManager | t.CPtr, level: t.CInt, wbits: t.CInt, mem_level: t.CInt, strategy: t.CInt) -> zdeflate_stream | t.CPtr: n: t.CSizeT = zdeflate_stream.__sizeof__() s: zdeflate_stream | t.CPtr = zdef.zdef_alloc(pool, n) if not s: return None memset(s, 0, zdeflate_stream.__sizeof__()) s.pool = pool s.writer = zdef.zbit_writer(pool) s.window = BYTEPTR(zdef.zdef_alloc(pool, zdef.ZDEFLATE_WINDOW_SIZE)) s.hash_head = INTPTR(zdef.zdef_alloc(pool, zdef.ZDEFLATE_HASH_SIZE * int.__sizeof__())) s.hash_prev = INTPTR(zdef.zdef_alloc(pool, zdef.ZDEFLATE_WINDOW_SIZE * int.__sizeof__())) if not s.window or not s.hash_head or not s.hash_prev: s.destroy() return None memset(s.hash_head, -1, zdef.ZDEFLATE_HASH_SIZE * int.__sizeof__()) memset(s.hash_prev, -1, zdef.ZDEFLATE_WINDOW_SIZE * int.__sizeof__()) s.level = level s.wbits = wbits s.strategy = strategy s.is_finished = 0 s.adler = 1 if wbits > 0: s.writer.write_zlib_header(wbits, level) elif wbits == 0: s.writer.write_zlib_header(15, level) return s def zdeflate_one_shot(pool: memhub.MemManager | t.CPtr, data: UINT8PTR, length: t.CSizeT, level: t.CInt, wbits: t.CInt, out_len: t.CSizeT | t.CPtr) -> t.CUInt8T | t.CPtr: s: zdeflate_stream | t.CPtr = zdeflate_create(pool, level, wbits, 8, 0) if not s: return None memset(s.hash_head, -1, zdef.ZDEFLATE_HASH_SIZE * int.__sizeof__()) memset(s.hash_prev, -1, zdef.ZDEFLATE_WINDOW_SIZE * int.__sizeof__()) s.adler = zchecksum.zchecksum_adler32(data, length, 1) if length == 0: s.writer.write_bits(1, 1) s.writer.write_bits(1, 2) lit_tree = zhuff.zhuff_tree() lit_tree.build_fixed_lit_tree() lit_tree.encode_symbol(256, c.Addr(s.writer)) else: s.compress_block(data, length, 1) s.is_finished = 1 if wbits >= 0: s.writer.align() adler: t.CUInt32T = s.adler s.writer.buf[s.writer.byte_pos + 0] = (adler >> 24) & 0xFF s.writer.buf[s.writer.byte_pos + 1] = (adler >> 16) & 0xFF s.writer.buf[s.writer.byte_pos + 2] = (adler >> 8) & 0xFF s.writer.buf[s.writer.byte_pos + 3] = (adler >> 0) & 0xFF s.writer.byte_pos += 4 c.Set(c.Deref(out_len), s.writer.total()) result: UINT8PTR = UINT8PTR(zdef.zdef_alloc(pool, c.Deref(out_len))) if result: memcpy(result, s.writer.buf, c.Deref(out_len)) s.destroy() return result