import mpq


class MpqFile:
    def __init__(self, name):
        self.map = open(name, "r+b")
        self.m = mpq.Mpq(self.map)

    def open_file(self, name):
        return self.m.find_file(name)

    def deprotect_map(self):
        bt = self.m.raw_bytes[self.m.HOffset + self.m.header.BTOffset: self.m.HOffset + self.m.header.BTOffset + self.m.header.BTEntries * 16]
        self.m.header.BTOffset = self.m.header.HTOffset + self.m.header.HTEntries * 16
        bh = b"MPQ\x1A" + self.m.split_to_bytes([32, self.m.size, self.m.header.SectorSizeShift << 16]) \
            + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries, self.m.header.BTEntries])
        # print(list(bh))
        self.m.raw_bytes = self.m.raw_bytes[:0x200] + bh + self.m.raw_bytes[self.m.HOffset + 32:] + bt
        # self.m.raw_bytes = self.m.raw_bytes[:self.m.HOffset] + bh + self.m.raw_bytes[self.m.HOffset + 32:]

    def protect_map(self):
        raw = b""   # self.m.raw_bytes[:0x200]
        bt = []
        # offsets = []
        # listfile = self.m.find_file("(listfile)").decode().split("\r\n") + ["(listfile)"]
        # print(listfile)

        print(len(self.m.block_table))
        for i in range(len(self.m.block_table)):
            # for j in listfile:
            #     if self.m.is_file_exist(j) == i:
            #        # print(j)
            #        name = j
            #        break
            file = self.m.find_file("", i)
            # print(i, file)
            # if i <= 8:
            # key = self.m.hash_string(name, self.m.MPQ_HASH_FILE_KEY)
            # key = (key + self.m.block_table[n].filePos) ^ self.m.block_table[n].f_size

            c_file = self.m.compress_file(file)
            raw = c_file + raw  # raw += c_file
            bt += [-len(raw), len(c_file), len(file), self.m.MPQ_FILE_EXISTS | self.m.MPQ_FILE_COMPRESS]
            # else:
            #    bt += [len(raw), len(file), len(file), self.m.MPQ_FILE_EXISTS]
            #    raw += file
            # offsets.append(len(raw))

        l = len(raw)
        ht = self.m.raw_bytes[self.m.HOffset + self.m.header.HTOffset: self.m.HOffset + self.m.header.HTOffset + self.m.header.HTEntries * 16]
        raw = ht + raw
        # + len(bt) * 4
        # for i in range(len(bt) // 4):
        #    bt[i * 4] -= l
        self.m.encrypt_mpq_block(bt, len(bt), self.m.MPQ_KEY_BLOCK_TABLE)
        bt = self.m.split_to_bytes(bt)
        # raw = bt + raw
        raw = bt + raw
        self.m.header.HTOffset = -l - len(ht)  # -len(ht) - len(bt)
        self.m.header.BTOffset = -l - len(ht) - len(bt)  # -len(bt)
        # self.m.size = len(raw) + 544

        # bh = b"MPQ\x1A" + self.m.split_to_bytes([self.m.header.HeaderSize, self.m.size]) \
        #                + self.m.split_to_bytes([self.m.header.FormatVersion, self.m.header.SectorSizeShift], b"", 2) \
        #              + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries,
        #                                    self.m.header.BTEntries])

        # bh = b"MPQ\x1A" + b"\x11" * 10 + self.m.split_to_bytes([self.m.header.SectorSizeShift], b"", 2) \
        bh = b"MPQ\x1A" + b"\xFF" * 10 + bytes([self.m.header.SectorSizeShift, 0]) \
            + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries, self.m.header.BTEntries])

        # print(len(bh))
        # print(len(raw) % 0x200)
        i = 0x200 - len(raw) % 0x200
        # print(self.m.raw_bytes[:i])
        # raw = self.m.raw_bytes[:0x200] + bytes([0] * i) + raw + bh
        raw = self.m.raw_bytes[:i] + raw + bh
        self.m.raw_bytes = raw[:]

    def recreate_map(self):
        raw = b""   # self.m.raw_bytes[:0x200]
        bt = []
        # offsets = []
        # listfile = self.m.find_file("(listfile)").decode().split("\r\n") + ["(listfile)"]
        # print(listfile)
        # self.m.block_table = self.m.block_table[:2]
        # print(self.m.block_table[1616].filePos)
        for i in range(len(self.m.block_table)):
            # for j in listfile:
            #     if self.m.is_file_exist(j) == i:
            #        # print(j)
            #        name = j
            #        break
            try:
                file = self.m.find_file("", i)
            except OSError:
                print(i)
            # print(i, file)
            # if i <= 8:
            # key = self.m.hash_string(name, self.m.MPQ_HASH_FILE_KEY)
            # key = (key + self.m.block_table[n].filePos) ^ self.m.block_table[n].f_size

            if not file:
                # c_file = b""
                print(i)
                bt += [20 + len(raw), 0, 0, self.m.MPQ_FILE_EXISTS | self.m.MPQ_FILE_COMPRESS]
            else:
                c_file = self.m.compress_file(file)
                bt += [32 + len(raw), len(c_file), len(file), self.m.MPQ_FILE_EXISTS | self.m.MPQ_FILE_COMPRESS]
                raw += c_file
            # else:
            #    bt += [len(raw), len(file), len(file), self.m.MPQ_FILE_EXISTS]
            #    raw += file
            # offsets.append(len(raw))

        # print(i)
        l = len(raw)
        ht = self.m.raw_bytes[self.m.HOffset + self.m.header.HTOffset: self.m.HOffset + self.m.header.HTOffset + self.m.header.HTEntries * 16]
        raw += ht
        # + len(bt) * 4
        # for i in range(len(bt) // 4):
        #    bt[i * 4] -= l
        self.m.encrypt_mpq_block(bt, len(bt), self.m.MPQ_KEY_BLOCK_TABLE)
        bt = self.m.split_to_bytes(bt)
        # raw = bt + raw
        raw += bt
        self.m.header.HTOffset = l + 32  # -len(ht) - len(bt)
        self.m.header.BTOffset = l + len(ht) + 32  # -len(bt)
        # self.m.size = len(raw) + 544

        # bh = b"MPQ\x1A" + self.m.split_to_bytes([self.m.header.HeaderSize, self.m.size]) \
        #                + self.m.split_to_bytes([self.m.header.FormatVersion, self.m.header.SectorSizeShift], b"", 2) \
        #              + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries,
        #                                    self.m.header.BTEntries])

        # bh = b"MPQ\x1A" + b"\x11" * 10 + self.m.split_to_bytes([self.m.header.SectorSizeShift], b"", 2) \
        print(len(bt))
        bh = b"MPQ\x1A" + b"\xFF" * 8 + b"00" + bytes([self.m.header.SectorSizeShift, 0]) \
            + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries, len(bt) // 16])

        # print(len(bh))
        # print(len(raw) % 0x200)
        i = 0x200  # - len(raw) % 0x200
        # print(self.m.raw_bytes[:i])
        # raw = self.m.raw_bytes[:0x200] + bytes([0] * i) + raw + bh
        raw = self.m.raw_bytes[:i] + bh + raw
        self.m.raw_bytes = raw[:]

    @staticmethod
    def calc_ot_len(file_size, sector_size):
        l = file_size // sector_size
        if file_size % sector_size > 0:
            l += 1
        return l

    def encrypt_file(self, file, name):
        file_size = len(file)
        sector_size = 512 * 2 ** self.m.header.SectorSizeShift
        ot_len = self.calc_ot_len(file_size, sector_size)
        ot_size = ot_len * 4 + 4
        print(len(file), sector_size, ot_len)
        ot = [ot_size + i * sector_size for i in range(ot_len)] + [ot_size + file_size]
        print(ot)

        encrypted_file = b""    # ot[:]
        key = self.m.hash_string(name, self.m.MPQ_HASH_FILE_KEY) - 1
        # print(key)
        # self.m.encrypt_mpq_block(encrypted_file, len(encrypted_file), key)
        # encrypted_file = self.m.split_to_bytes(encrypted_file)

        for i in range(1, len(ot)):
            # print(i)
            sector = file[ot[i - 1] - ot_size: ot[i] - ot_size]
            # print(sector)
            last = sector[len(sector) // 4 * 4:]
            # print(last)
            sector = [int.from_bytes(sector[j * 4: j * 4 + 4], "little") for j in range(len(sector) // 4)]
            self.m.encrypt_mpq_block(sector, len(sector), key + i)

            sector = self.m.split_to_bytes(sector, last)
            encrypted_file += sector

        return encrypted_file

    def decrypt_file(self, name):
        # todo
        pass

    def delete_file(self, name):
        raw = self.m.raw_bytes

        if isinstance(name, int):
            n = name
        else:
            n = self.m.is_file_exist(name)
        # print(n)
        # old_header = raw[self.m.HOffset: self.m.HOffset + 32]
        # print(self.m.size)
        # print(self.m.header.ArchiveSize)

        # n = self.m.is_file_exist("war3map.j")
        file_offset = self.m.block_table[n].filePos
        old_size = self.m.block_table[n].c_size
        raw = raw[:self.m.HOffset + file_offset] + raw[self.m.HOffset + file_offset + old_size:]
        d_size = -old_size
        # self.m.blockTable[n].flags = 0x80000000

        bt = []
        self.m.header.BTEntries -= 1
        del self.m.block_table[n]
        for i in self.m.block_table:
            if i.filePos > file_offset:
                i.filePos += d_size
            bt += [i.filePos, i.c_size, i.f_size, i.flags]
        #print(bt[5*4:6*4])
        #for i in range(self.m.header.BTEntries):
        #    print([hex(j) for j in bt[i * 4: i * 4 + 4]])

        self.m.encrypt_mpq_block(bt, len(bt), self.m.MPQ_KEY_BLOCK_TABLE)
        bt = self.m.split_to_bytes(bt)

        #print(name)
        ht = self.m.raw_ht[:]
        #for i in range(self.m.header.HTEntries):
        #    print([hex(j) for j in ht[i * 4: i * 4 + 4]])
        # print([hex(i) for i in ht[4 * 16: 4 * 17]])
        hti = [ht[i * 4 + 3] for i in range(self.m.header.HTEntries)].index(n)
        ht[hti * 4: hti * 4 + 4] = [0xFFFFFFFF] * 3 + [0xFFFFFFFE]
        for i in range(self.m.header.HTEntries):
            if 0xFFFFFFFE > ht[i * 4 + 3] > n:
                self.m.hash_table[ht[i * 4]][ht[i * 4 + 1]] -= 1
                ht[i * 4 + 3] -= 1

        # print(ht[38 * 4: 39 * 4])
        self.m.raw_ht = ht[:]

        self.m.encrypt_mpq_block(ht, len(ht), self.m.MPQ_KEY_HASH_TABLE)
        ht = self.m.split_to_bytes(ht)

        if self.m.header.HTOffset > file_offset:
            self.m.header.HTOffset += d_size
        if self.m.header.BTOffset > file_offset:
            self.m.header.BTOffset += d_size

        bh = b"MPQ\x1A\x20\x00\x00\x00" + self.m.split_to_bytes([self.m.size]) \
                                        + self.m.split_to_bytes([self.m.header.FormatVersion, self.m.header.SectorSizeShift], b"", 2) \
                                        + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries,
                                                            self.m.header.BTEntries])

        # print(bh)
        # print(old_header)
        # print(len(bh))
        raw = raw[:self.m.HOffset] + bh + raw[self.m.HOffset + 32:]
        raw = raw[:self.m.HOffset + self.m.header.HTOffset] + ht + raw[self.m.HOffset + self.m.header.HTOffset + self.m.header.HTEntries * 16:]
        raw = raw[:self.m.HOffset + self.m.header.BTOffset] + bt + raw[self.m.HOffset + self.m.header.BTOffset + self.m.header.BTEntries * 16:]
        self.m.raw_bytes = raw

    def edit_file(self, name, new_file):
        raw = self.m.raw_bytes

        n = self.m.is_file_exist(name)
        # old_header = raw[self.m.HOffset: self.m.HOffset + 32]
        # print(self.m.size)
        # print(self.m.header.ArchiveSize)

        # n = self.m.is_file_exist("war3map.j")
        file_offset = self.m.block_table[n].filePos
        old_size = self.m.block_table[n].c_size
        # new_file = encrypt_file(self.m.find_file(name), name)
        # key = self.m.hash_string(name, self.m.MPQ_HASH_FILE_KEY)
        # key = (key + self.m.block_table[n].filePos) ^ self.m.block_table[n].f_size
        new_f_size = len(new_file)
        new_file = self.m.compress_file(new_file, False)  # True, key - 1)
        new_c_size = len(new_file)
        raw = raw[:self.m.HOffset + file_offset] + new_file + raw[self.m.HOffset + file_offset + old_size:]
        d_size = new_c_size - old_size
        # print(hex(d_size))

        bt = []
        self.m.block_table[n].c_size = new_c_size
        self.m.block_table[n].f_size = new_f_size
        self.m.block_table[n].flags = self.m.MPQ_FILE_EXISTS | self.m.MPQ_FILE_COMPRESS  # | self.m.MPQ_FILE_ENCRYPTED | self.m.MPQ_FILE_FIX_KEY
        for i in self.m.block_table:
            if i.filePos > file_offset:
                i.filePos += d_size
            # if self.m.block_table[n].flags & self.m.MPQ_FILE_FIX_KEY:
            #    pass
            bt += [i.filePos, i.c_size, i.f_size, i.flags]
        # print(bt[5*4:6*4])
        # for i in range(self.m.header.BTEntries):
        #    print([hex(j) for j in bt[i * 4: i * 4 + 4]])
        self.m.encrypt_mpq_block(bt, len(bt), self.m.MPQ_KEY_BLOCK_TABLE)
        bt = self.m.split_to_bytes(bt)

        if self.m.header.HTOffset > file_offset:
            self.m.header.HTOffset += d_size
        if self.m.header.BTOffset > file_offset:
            self.m.header.BTOffset += d_size
        self.m.size += d_size

        bh = b"MPQ\x1A\x20\x00\x00\x00" + self.m.split_to_bytes([self.m.size]) \
                                        + self.m.split_to_bytes([self.m.header.FormatVersion, self.m.header.SectorSizeShift], b"", 2) \
                                        + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries,
                                                            self.m.header.BTEntries])

        # print(bh)
        # print(old_header)
        # print(len(bh))
        raw = raw[:self.m.HOffset] + bh + raw[self.m.HOffset + 32:]
        # raw = raw[:self.m.HOffset + self.m.header.HTOffset] + ht + raw[self.m.HOffset + self.m.header.HTOffset + self.m.header.HTEntries * 16:]
        raw = raw[:self.m.HOffset + self.m.header.BTOffset] + bt + raw[self.m.HOffset + self.m.header.BTOffset + self.m.header.BTEntries * 16:]
        self.m.raw_bytes = raw

    def add_file(self, name, new_file):
        raw = self.m.raw_bytes
        offset = self.m.HOffset + self.m.header.HTOffset
        f_size = len(new_file)
        new_file = self.m.compress_file(new_file, False)  # True, key - 1)
        c_size = len(new_file)
        flags = self.m.MPQ_FILE_EXISTS | self.m.MPQ_FILE_COMPRESS
        raw = raw[:offset] + new_file  # + raw[offset:]
        self.m.raw_bt += [self.m.header.HTOffset, c_size, f_size, flags]
        bt = self.m.raw_bt[:]
        self.m.encrypt_mpq_block(bt, len(bt), self.m.MPQ_KEY_BLOCK_TABLE)
        bt = self.m.split_to_bytes(bt)

        h1 = self.m.hash_string(name, self.m.MPQ_HASH_TABLE_INDEX) % self.m.header.HTEntries
        h2 = self.m.hash_string(name, self.m.MPQ_HASH_NAME_A)
        h3 = self.m.hash_string(name, self.m.MPQ_HASH_NAME_B)
        while self.m.raw_ht[h1 * 4 + 3] != 0xFFFFFFFF:
            h1 = (h1 + 1) % self.m.header.HTEntries
        self.m.raw_ht = self.m.raw_ht[:4 * h1:] + [h2, h3, 0, self.m.header.BTEntries] + self.m.raw_ht[4 * h1 + 4:]
        ht = self.m.raw_ht[:]
        self.m.encrypt_mpq_block(ht, len(ht), self.m.MPQ_KEY_HASH_TABLE)
        ht = self.m.split_to_bytes(ht)

        self.m.header.HTOffset += c_size
        self.m.header.BTOffset += c_size
        self.m.header.BTEntries += 1
        bh = b"MPQ\x1A" + self.m.split_to_bytes([32, self.m.size, self.m.header.SectorSizeShift << 16]) \
            + self.m.split_to_bytes([self.m.header.HTOffset, self.m.header.BTOffset, self.m.header.HTEntries, self.m.header.BTEntries])

        raw = raw[:0x200] + bh + raw[0x220: self.m.HOffset + self.m.header.HTOffset] + ht + bt
        self.m.raw_bytes = raw

    def save_archive(self):
        self.map.seek(0)
        self.map.truncate()
        self.map.write(self.m.raw_bytes)
        self.map.close()

