Files
TransPyC/includes/zlib/zhuff.py
2026-07-18 19:25:40 +08:00

357 lines
13 KiB
Python

#include <string.h>
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