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

457 lines
18 KiB
Python

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