357 lines
13 KiB
Python
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
|