232 lines
9.8 KiB
Python
232 lines
9.8 KiB
Python
from io import BytesIO
|
|
from collections import defaultdict
|
|
|
|
class LecteurDeBits:
|
|
|
|
def __init__(self, entree : bytes):
|
|
self.octets = BytesIO(entree)
|
|
self.bits_non_lus = 0
|
|
self.taille_bits_non_lus = 0
|
|
|
|
def lire_bits(self, nombre_bits : int) -> int:
|
|
while self.taille_bits_non_lus < nombre_bits:
|
|
prochain_octet = self.octets.read(1)
|
|
if not prochain_octet:
|
|
raise EOFError
|
|
self.bits_non_lus |= prochain_octet[0] << self.taille_bits_non_lus
|
|
self.taille_bits_non_lus += 8
|
|
masque_bits_lus = (1 << nombre_bits) - 1
|
|
bits_lus = self.bits_non_lus & masque_bits_lus
|
|
self.bits_non_lus >>= nombre_bits
|
|
self.taille_bits_non_lus -= nombre_bits
|
|
return bits_lus
|
|
|
|
def lire_octets(self, nombre_octets : int) -> bytes:
|
|
self.aligner_bits_sur_octet()
|
|
octets_lus = self.octets.read(nombre_octets)
|
|
if len(octets_lus) < nombre_octets:
|
|
raise EOFError
|
|
return octets_lus
|
|
|
|
def aligner_bits_sur_octet(self):
|
|
|
|
self.bits_non_lus = 0
|
|
self.taille_bits_non_lus = 0
|
|
|
|
class LecteurDeBitsRangeCode:
|
|
|
|
def __init__(self, entree : bytes):
|
|
self.octets = BytesIO(entree)
|
|
self.taille_code = 0xffffffff
|
|
self.code = 0
|
|
assert self.octets.read(1)[0] == 0
|
|
|
|
for position in range(4):
|
|
self.code = (self.code << 8) | self.octets.read(1)[0]
|
|
|
|
assert self.code < self.taille_code
|
|
|
|
def reprendre_des_bits_si_besoin(self):
|
|
if self.taille_code <= 0xffffff:
|
|
self.taille_code <<= 8
|
|
self.code <<= 8
|
|
self.code |= self.octets.read(1)[0]
|
|
|
|
class RangeDecoder:
|
|
|
|
def __init__(self, lecteur_de_bits : LecteurDeBitsRangeCode, probabilite_initiale : int = 0x400):
|
|
self.lecteur_de_bits : LecteurDeBitsRangeCode = lecteur_de_bits
|
|
self.probabilite : int = probabilite_initiale # 0 - 0x800
|
|
|
|
def lire_bit(self, utiliser_probas = True):
|
|
if utiliser_probas:
|
|
milieu_du_code_pondere = self.lecteur_de_bits.taille_code // 0x800 * self.probabilite
|
|
else:
|
|
milieu_du_code_pondere = self.lecteur_de_bits.taille_code // 2
|
|
|
|
if self.lecteur_de_bits.code < milieu_du_code_pondere:
|
|
bit_lu = 0
|
|
self.probabilite += (0x800 - self.probabilite) // 32
|
|
self.lecteur_de_bits.taille_code = milieu_du_code_pondere
|
|
else:
|
|
bit_lu = 1
|
|
self.probabilite -= self.probabilite // 32
|
|
self.lecteur_de_bits.code -= milieu_du_code_pondere
|
|
if utiliser_probas:
|
|
self.lecteur_de_bits.taille_code -= milieu_du_code_pondere
|
|
else:
|
|
self.lecteur_de_bits.taille_code = milieu_du_code_pondere
|
|
self.lecteur_de_bits.reprendre_des_bits_si_besoin()
|
|
return bit_lu
|
|
|
|
class DecodeurLZMASPD:
|
|
|
|
def decode(self, entree : bytes) -> bytes:
|
|
self.flux_decompresse = b''
|
|
self.dernieres_distances : List[int] = [0] * 4
|
|
lecteur_de_bits = LecteurDeBits(entree)
|
|
properties = lecteur_de_bits.lire_bits(8)
|
|
self.literal_context_bits = properties % 9
|
|
literal_position_bits = (properties // 9) % 5
|
|
position_bits = properties // 9 // 5
|
|
|
|
if position_bits > 4:
|
|
raise ValueError('LZMA invalid')
|
|
|
|
self.taille_fenetre = lecteur_de_bits.lire_bits(32)
|
|
self.uncompressed_size = lecteur_de_bits.lire_bits(64)
|
|
lecteur_de_bits = LecteurDeBitsRangeCode(entree[13:])
|
|
self.nom_vers_range_decodeur : Dict[tuple, RangeDecoder] = defaultdict(lambda: RangeDecoder(lecteur_de_bits))
|
|
self.state = 0
|
|
self.masque_pos_state = (1 << position_bits) - 1
|
|
self.masque_lit_state = (1 << literal_position_bits) - 1
|
|
|
|
while len(self.flux_decompresse) < self.uncompressed_size:
|
|
try:
|
|
pos_state = len(self.flux_decompresse) & self.masque_pos_state
|
|
bit_choice = self.nom_vers_range_decodeur[('IsMatch', self.state, pos_state)].lire_bit()
|
|
if bit_choice == 0:
|
|
self.LITERAL()
|
|
elif bit_choice == 1:
|
|
is_rep = self.nom_vers_range_decodeur[('IsRep', self.state)].lire_bit()
|
|
if is_rep == 0:
|
|
self.MATCH()
|
|
elif is_rep == 1:
|
|
is_rep0 = self.nom_vers_range_decodeur[('IsRepG0', self.state)].lire_bit()
|
|
if is_rep0 == 0:
|
|
is_rep0_long = self.nom_vers_range_decodeur[('IsRep0Long', self.state, pos_state)].lire_bit()
|
|
if is_rep0_long == 0:
|
|
self.SHORTREP()
|
|
elif is_rep0_long == 1:
|
|
self.LONGREP(0)
|
|
elif is_rep0 == 1:
|
|
is_rep1 = self.nom_vers_range_decodeur[('IsRepG1', self.state)].lire_bit()
|
|
if is_rep1 == 0:
|
|
self.LONGREP(1)
|
|
elif is_rep1 == 1:
|
|
is_rep2 = self.nom_vers_range_decodeur[('IsRepG2', self.state)].lire_bit()
|
|
if is_rep2 == 0:
|
|
self.LONGREP(2)
|
|
elif is_rep2 == 1:
|
|
self.LONGREP(3)
|
|
except EOFError:
|
|
break
|
|
|
|
return self.flux_decompresse
|
|
|
|
def LITERAL(self):
|
|
dernier_octet_decompresse = self.flux_decompresse[-1] if self.flux_decompresse else 0
|
|
octet_lu = self.bit_tree_decode(('LiteralNormal',
|
|
len(self.flux_decompresse) & self.masque_lit_state, # total_pos
|
|
dernier_octet_decompresse >> (8 - self.literal_context_bits), # prev_byte
|
|
), None, 8, use_pos_state = False)
|
|
self.flux_decompresse += bytes([octet_lu])
|
|
if self.state > 9:
|
|
self.state -= 6
|
|
elif self.state > 3:
|
|
self.state -= 3
|
|
else:
|
|
self.state = 0
|
|
|
|
def MATCH(self):
|
|
match_len = 2 + self.len_decode('LenDecoder')
|
|
pos_slot = self.bit_tree_decode(('PosSlot', min(5, match_len)), None, 6, use_pos_state = False)
|
|
if pos_slot >= 4:
|
|
num_direct_bits = (pos_slot >> 1) - 1
|
|
distance = (2 | (pos_slot & 1)) << num_direct_bits
|
|
if pos_slot < 14:
|
|
distance += self.bit_tree_decode('SpecPos', None, num_direct_bits + (distance - pos_slot - 1),
|
|
use_pos_state = False, reverse = True,
|
|
debut_bit_tree = distance - pos_slot - 1)
|
|
else:
|
|
distance += self.bit_tree_decode('AlignFixed', None, num_direct_bits - 4, utiliser_probas = False) << 4
|
|
distance += self.bit_tree_decode('Align', None, 4, use_pos_state = False, reverse = True)
|
|
else:
|
|
distance = pos_slot
|
|
|
|
self.dernieres_distances.append(distance)
|
|
if distance == 0xffffffff:
|
|
raise EOFError
|
|
assert distance < len(self.flux_decompresse)
|
|
assert distance < self.taille_fenetre
|
|
self.repeter_donnees(distance, match_len)
|
|
if self.state < 7:
|
|
self.state = 7
|
|
else:
|
|
self.state = 10
|
|
|
|
def SHORTREP(self): # Réutiliser la dernière distance pour un octet
|
|
if self.state < 7:
|
|
self.state = 9
|
|
else:
|
|
self.state = 11
|
|
self.repeter_donnees(self.dernieres_distances[-1], 1)
|
|
|
|
def LONGREP(self, num): # Réutiliser l'une des dernières distances pour une taille donnée
|
|
match_len = 2 + self.len_decode('RepLenDecoder')
|
|
self.dernieres_distances.append(self.dernieres_distances.pop(-(1 + num)))
|
|
distance = self.dernieres_distances[-1]
|
|
self.repeter_donnees(distance, match_len)
|
|
if self.state < 7:
|
|
self.state = 8
|
|
else:
|
|
self.state = 11
|
|
|
|
def repeter_donnees(self, distance, match_len):
|
|
debut_slice = len(self.flux_decompresse) - (1 + distance)
|
|
fin_slice = match_len
|
|
a_repeter = self.flux_decompresse[debut_slice:debut_slice + fin_slice]
|
|
fin_slice -= len(self.flux_decompresse) - debut_slice
|
|
while fin_slice > 0:
|
|
a_repeter += self.flux_decompresse[debut_slice:debut_slice + fin_slice]
|
|
fin_slice -= min(len(self.flux_decompresse), debut_slice + fin_slice) - debut_slice
|
|
|
|
self.flux_decompresse += a_repeter
|
|
|
|
def len_decode(self, len_decoder_name):
|
|
if self.nom_vers_range_decodeur[('LenChoice', len_decoder_name)].lire_bit() == 0:
|
|
return self.bit_tree_decode('LenLow', len_decoder_name, 3)
|
|
else:
|
|
if self.nom_vers_range_decodeur[('LenChoice2', len_decoder_name)].lire_bit() == 0:
|
|
return (1 << 3) + self.bit_tree_decode('LenMid', len_decoder_name, 3)
|
|
else:
|
|
return (1 << 4) + self.bit_tree_decode('LenHigh', len_decoder_name, 8, use_pos_state = False)
|
|
|
|
def bit_tree_decode(self, bit_tree_decoder_name, len_decoder_name, num_bits,
|
|
use_pos_state = True, utiliser_probas = True, reverse = False,
|
|
bit_tree_lu = 0, debut_bit_tree = 0):
|
|
for position_bit in range(debut_bit_tree, num_bits):
|
|
bit_lu = self.nom_vers_range_decodeur[(
|
|
bit_tree_decoder_name,
|
|
len_decoder_name,
|
|
bit_tree_lu,
|
|
position_bit,
|
|
(len(self.flux_decompresse) & self.masque_pos_state) if use_pos_state else None,
|
|
)].lire_bit(utiliser_probas = utiliser_probas)
|
|
if not reverse:
|
|
bit_tree_lu <<= 1
|
|
bit_tree_lu |= bit_lu
|
|
else:
|
|
bit_tree_lu |= bit_lu << (position_bit - debut_bit_tree)
|
|
return bit_tree_lu
|