#include from stdint import * import zlib.zdef as zdef import memhub import t, c ZHUFF_MAX_CODES: t.CDefine = 288 ZHUFF_MAX_BITS: t.CDefine = 15 class zhuff_code: code: UINT bits: t.CInt @t.Object class zhuff_tree: pool: memhub.MemManager | t.CPtr codes: t.CArray[zhuff_code, ZHUFF_MAX_CODES] count: t.CInt max_bits: t.CInt def __init__(self): self.pool = None def build_codes(self, freqs: INTPTR, count: t.CInt, max_bits: t.CInt): lengths: t.CArray[t.CInt, ZHUFF_MAX_CODES] bl_count: t.CArray[t.CInt, ZHUFF_MAX_BITS + 1] next_code: t.CArray[UINT, ZHUFF_MAX_BITS + 1] self.count = count self.max_bits = max_bits self.build_code_lengths(lengths, freqs, count, max_bits) memset(bl_count, 0, bl_count.__sizeof__()) for i in range(count): if lengths[i] > 0: bl_count[lengths[i]] += 1 code: t.CUnsignedInt = 0 next_code[0] = 0 for bits in range(1, max_bits + 1): code = (code + bl_count[bits - 1]) << 1 next_code[bits] = code for i in range(count): self.codes[i].bits = lengths[i] if lengths[i] > 0: self.codes[i].code = next_code[lengths[i]] next_code[lengths[i]] += 1 else: self.codes[i].code = 0 def build_fixed_lit_tree(self): lengths: t.CArray[t.CInt, 288] bl_count: t.CArray[t.CInt, 16] next_code: t.CArray[t.CUnsignedInt, 16] self.get_fixed_lit_lengths(lengths) memset(bl_count, 0, bl_count.__sizeof__()) for i in range(288): if lengths[i] > 0: bl_count[lengths[i]] += 1 code: t.CUnsignedInt = 0 next_code[0] = 0 for bits in range(1, 9 + 1): code = (code + bl_count[bits - 1]) << 1 next_code[bits] = code self.count = 288 self.max_bits = 9 for i in range(288): self.codes[i].bits = lengths[i] if lengths[i] > 0: self.codes[i].code = next_code[lengths[i]] next_code[lengths[i]] += 1 else: self.codes[i].code = 0 def build_fixed_dist_tree(self): lengths: t.CArray[t.CInt, 32] bl_count: t.CArray[t.CInt, 16] next_code: t.CArray[t.CUnsignedInt, 16] self.get_fixed_dist_lengths(lengths) memset(bl_count, 0, bl_count.__sizeof__()) for i in range(32): if lengths[i] > 0: bl_count[lengths[i]] += 1 code: t.CUnsignedInt = 0 next_code[0] = 0 for bits in range(1, 5 + 1): code = (code + bl_count[bits - 1]) << 1 next_code[bits] = code self.count = 32 self.max_bits = 5 for i in range(32): self.codes[i].bits = lengths[i] if lengths[i] > 0: self.codes[i].code = next_code[lengths[i]] next_code[lengths[i]] += 1 else: self.codes[i].code = 0 def encode_symbol(self, symbol: int, writer: zdef.zbit_writer | t.CPtr): if symbol < 0 or symbol >= self.count: return if self.codes[symbol].bits == 0: return writer.write_bits_rev(self.codes[symbol].code, self.codes[symbol].bits) def build_code_lengths(self, lengths: t.CInt | t.CPtr, freqs: t.CInt | t.CPtr, count: t.CInt, max_bits: t.CInt): bl_count: t.CArray[t.CInt, ZHUFF_MAX_BITS + 2] sort_count: t.CInt = 0 memset(bl_count, 0, bl_count.__sizeof__()) memset(lengths, 0, int.__sizeof__() * count) for i in range(count): if freqs[i] > 0: sort_count += 1 if sort_count == 0: return if sort_count == 1: for i in range(count): if freqs[i] > 0: lengths[i] = 1 return items: t.CInt | t.CPtr = INTPTR(zdef.zdef_alloc(self.pool, sort_count * int.__sizeof__())) item_freqs: t.CInt | t.CPtr = INTPTR(zdef.zdef_alloc(self.pool, sort_count * int.__sizeof__())) idx: t.CInt = 0 for i in range(count): if freqs[i] > 0: items[idx] = i item_freqs[idx] = freqs[i] idx += 1 for i in range(sort_count): key_item: t.CInt = items[i] key_freq: t.CInt = item_freqs[i] j: t.CInt = i - 1 while j >= 0 and item_freqs[j] > key_freq: items[j + 1] = items[j] item_freqs[j + 1] = item_freqs[j] j -= 1 items[j + 1] = key_item item_freqs[j + 1] = key_freq parent: INTPTR = INTPTR(zdef.zdef_alloc(self.pool, (sort_count * 2) * int.__sizeof__())) for i in range(sort_count * 2): parent[i] = -1 heap: INTPTR = INTPTR(zdef.zdef_alloc(self.pool, (sort_count + 1) * int.__sizeof__())) heap_size: t.CInt = 0 for i in range(sort_count): heap[heap_size] = i heap_size += 1 pos: t.CInt = heap_size - 1 while pos > 0: par: t.CInt = (pos - 1) / 2 if item_freqs[heap[par]] > item_freqs[heap[pos]]: tmp: t.CInt = heap[par] heap[par] = heap[pos] heap[pos] = tmp pos = par else: break combined_freq: INTPTR = INTPTR(zdef.zdef_alloc(self.pool, (sort_count * 2) * int.__sizeof__())) for i in range(sort_count): combined_freq[i] = item_freqs[i] internal: t.CInt = sort_count while heap_size > 1: a: t.CInt = heap[0] heap[0] = heap[heap_size - 1] heap_size -= 1 pos: t.CInt = 0 while True: left: t.CInt = 2 * pos + 1 right: t.CInt = 2 * pos + 2 smallest: t.CInt = pos if left < heap_size and combined_freq[heap[left]] < combined_freq[heap[smallest]]: smallest = left if right < heap_size and combined_freq[heap[right]] < combined_freq[heap[smallest]]: smallest = right if smallest != pos: tmp: t.CInt = heap[pos] heap[pos] = heap[smallest] heap[smallest] = tmp pos = smallest else: break b: t.CInt = heap[0] heap[0] = heap[heap_size - 1] heap_size -= 1 pos = 0 while True: left: t.CInt = 2 * pos + 1 right: t.CInt = 2 * pos + 2 smallest: t.CInt = pos if left < heap_size and combined_freq[heap[left]] < combined_freq[heap[smallest]]: smallest = left if right < heap_size and combined_freq[heap[right]] < combined_freq[heap[smallest]]: smallest = right if smallest != pos: tmp: t.CInt = heap[pos] heap[pos] = heap[smallest] heap[smallest] = tmp pos = smallest else: break if internal >= sort_count * 2 - 1: break combined_freq[internal] = combined_freq[a] + combined_freq[b] parent[a] = internal parent[b] = internal heap[heap_size] = internal heap_size += 1 pos = heap_size - 1 while pos > 0: par: t.CInt = (pos - 1) / 2 if combined_freq[heap[par]] > combined_freq[heap[pos]]: tmp: t.CInt = heap[par] heap[par] = heap[pos] heap[pos] = tmp pos = par else: break internal += 1 for i in range(sort_count): node: t.CInt = i length: t.CInt = 0 while parent[node] >= 0: length += 1 node = parent[node] lengths[items[i]] = length memset(bl_count, 0, bl_count.__sizeof__()) for i in range(count): if lengths[i] > 0: bl_count[lengths[i]] += 1 overflow: t.CInt = 0 for i in range(count): if lengths[i] > max_bits: overflow += 1 bl_count[lengths[i]] -= 1 lengths[i] = max_bits bl_count[max_bits] += 1 while overflow > 0: bits: t.CInt = max_bits - 1 while bits > 0 and bl_count[bits] == 0: bits -= 1 if bits == 0: break bl_count[bits] -= 1 bl_count[bits + 1] += 2 bl_count[max_bits] -= 1 overflow -= 2 sorted_items: INTPTR = INTPTR(zdef.zdef_alloc(self.pool, sort_count * int.__sizeof__())) si: t.CInt = 0 for i in range(count): if freqs[i] > 0: sorted_items[si] = i si += 1 for i in range(si - 1): for j in range(i + 1, si): if lengths[sorted_items[i]] < lengths[sorted_items[j]]: tmp: t.CInt = sorted_items[i] sorted_items[i] = sorted_items[j] sorted_items[j] = tmp sidx: t.CInt = 0 for bits in range(max_bits, 1, -1): n: t.CInt = bl_count[bits] while n > 0 and sidx < si: lengths[sorted_items[sidx]] = bits sidx += 1 n -= 1 zdef.zdef_free(self.pool, sorted_items) zdef.zdef_free(self.pool, heap) zdef.zdef_free(self.pool, parent) zdef.zdef_free(self.pool, combined_freq) zdef.zdef_free(self.pool, item_freqs) zdef.zdef_free(self.pool, items) def get_fixed_lit_lengths(self, lengths: INTPTR): for i in range( 0, 143 + 1): lengths[i] = 8 for i in range(143 + 1, 255 + 1): lengths[i] = 9 for i in range(255 + 1, 279 + 1): lengths[i] = 7 for i in range(279 + 1, 287 + 1): lengths[i] = 8 def get_fixed_dist_lengths(self, lengths: INTPTR): for i in range(32): lengths[i] = 5 def build_tree_from_lengths(self, lengths: INTPTR, count: t.CInt, max_bits: t.CInt): bl_count: t.CArray[t.CInt, 16] next_code: t.CArray[t.CUnsignedInt, 16] memset(bl_count, 0, bl_count.__sizeof__()) for i in range(count): if lengths[i] > 0: bl_count[lengths[i]] += 1 code: UINT = 0 next_code[0] = 0 for bits in range(1, max_bits + 1): code = (code + bl_count[bits - 1]) << 1 next_code[bits] = code self.count = count self.max_bits = max_bits for i in range(count): self.codes[i].bits = lengths[i] if lengths[i] > 0: self.codes[i].code = next_code[lengths[i]] next_code[lengths[i]] += 1 else: self.codes[i].code = 0 class zhuff_decode_node: children: t.CArray[t.CInt, 2] symbol: t.CInt @t.Object class zhuff_decode_tree: pool: memhub.MemManager | t.CPtr nodes: t.CArray[zhuff_decode_node, 2 * ZHUFF_MAX_CODES] node_count: t.CInt root: t.CInt def __init__(self): self.pool = None def build_decode_tree(self, ht: zhuff_tree | t.CPtr): self.node_count = 1 self.root = 0 self.nodes[0].children[0] = -1 self.nodes[0].children[1] = -1 self.nodes[0].symbol = -1 for i in range(ht.count): if ht.codes[i].bits == 0: continue node: int = 0 for bit in range(ht.codes[i].bits - 1, -1, -1): dir: t.CInt = (ht.codes[i].code >> bit) & 1 if bit > 0: # Internal node: traverse or create if self.nodes[node].children[dir] == -1: new_node: int = self.node_count self.node_count += 1 self.nodes[new_node].children[0] = -1 self.nodes[new_node].children[1] = -1 self.nodes[new_node].symbol = -1 self.nodes[node].children[dir] = new_node node = self.nodes[node].children[dir] else: # Leaf: bit 0 - set symbol on the child if self.nodes[node].children[dir] == -1: leaf: int = self.node_count self.node_count += 1 self.nodes[leaf].children[0] = -1 self.nodes[leaf].children[1] = -1 self.nodes[leaf].symbol = i self.nodes[node].children[dir] = leaf else: self.nodes[self.nodes[node].children[dir]].symbol = i def decode_symbol(self, reader: zdef.zbit_reader | t.CPtr) -> t.CInt: node: int = self.root while self.nodes[node].symbol == -1: bit: t.CUnsignedInt if reader.read_bits(1, c.Addr(bit)) != 0: return -1 dir: t.CInt = t.CInt(bit) if self.nodes[node].children[dir] == -1: return -1 node = self.nodes[node].children[dir] return self.nodes[node].symbol