#!/usr/bin/env python3
"""
axml_decode.py — Minimal, dependency-free Android Binary XML (AXML) decoder.
Decodes a compiled AndroidManifest.xml (or any compiled res XML) back to text.
Usage: python3 axml_decode.py <AndroidManifest.xml>
"""
import struct
import sys

# Chunk types
RES_STRING_POOL_TYPE = 0x0001
RES_XML_TYPE = 0x0003
RES_XML_START_NAMESPACE_TYPE = 0x0100
RES_XML_END_NAMESPACE_TYPE = 0x0101
RES_XML_START_ELEMENT_TYPE = 0x0102
RES_XML_END_ELEMENT_TYPE = 0x0103
RES_XML_CDATA_TYPE = 0x0104
RES_XML_RESOURCE_MAP_TYPE = 0x0180

# Attribute value types
TYPE_REFERENCE = 0x01
TYPE_STRING = 0x03
TYPE_FLOAT = 0x04
TYPE_INT_DEC = 0x10
TYPE_INT_HEX = 0x11
TYPE_INT_BOOLEAN = 0x12


class StringPool:
    def __init__(self, data, off):
        (self.type, self.header_size, self.size) = struct.unpack_from('<HHI', data, off)
        (self.string_count, self.style_count, self.flags,
         self.strings_start, self.styles_start) = struct.unpack_from('<IIIII', data, off + 8)
        self.is_utf8 = (self.flags & (1 << 8)) != 0
        self.offsets = []
        p = off + 28
        for i in range(self.string_count):
            self.offsets.append(struct.unpack_from('<I', data, p)[0])
            p += 4
        self.strings_base = off + self.strings_start
        self.data = data
        self._cache = {}

    def get(self, idx):
        if idx == 0xFFFFFFFF or idx < 0 or idx >= self.string_count:
            return None
        if idx in self._cache:
            return self._cache[idx]
        base = self.strings_base + self.offsets[idx]
        d = self.data
        if self.is_utf8:
            # utf-8: char count (u16-ish), then byte count, then bytes
            p = base
            # skip char length
            if d[p] & 0x80:
                p += 2
            else:
                p += 1
            if d[p] & 0x80:
                blen = ((d[p] & 0x7F) << 8) | d[p + 1]
                p += 2
            else:
                blen = d[p]
                p += 1
            s = d[p:p + blen].decode('utf-8', errors='replace')
        else:
            p = base
            if d[p] | (d[p + 1] << 8) & 0x8000:
                clen = d[p] | (d[p + 1] << 8)
                if clen & 0x8000:
                    clen = ((clen & 0x7FFF) << 16) | (d[p + 2] | (d[p + 3] << 8))
                    p += 4
                else:
                    p += 2
            s = d[p:p + clen * 2].decode('utf-16-le', errors='replace')
        self._cache[idx] = s
        return s


def decode(path):
    with open(path, 'rb') as f:
        data = f.read()
    magic, hsize, fsize = struct.unpack_from('<HHI', data, 0)
    off = 8
    pool = None
    out = []
    indent = 0
    ns_map = {}

    def res_val(val_type, val_data, pool):
        if val_type == TYPE_STRING:
            return pool.get(val_data)
        if val_type == TYPE_INT_BOOLEAN:
            return 'true' if val_data != 0 else 'false'
        if val_type == TYPE_REFERENCE:
            return '@' + hex(val_data)
        if val_type == TYPE_INT_HEX:
            return hex(val_data)
        if val_type == TYPE_INT_DEC:
            return str(val_data if val_data < 0x80000000 else val_data - 0x100000000)
        return str(val_data)

    while off < len(data):
        if off + 8 > len(data):
            break
        ctype, chsize, csize = struct.unpack_from('<HHI', data, off)
        if ctype == RES_STRING_POOL_TYPE:
            pool = StringPool(data, off)
        elif ctype == RES_XML_START_ELEMENT_TYPE:
            ns_idx, name_idx = struct.unpack_from('<II', data, off + 8 + 8)
            attr_start, attr_size, attr_count = struct.unpack_from('<HHH', data, off + 8 + 16)
            name = pool.get(name_idx)
            attrs = []
            ap = off + 8 + 8 + attr_start
            for i in range(attr_count):
                a_ns, a_name, a_rawval, a_typedval = struct.unpack_from('<IIII', data, ap)
                a_type = (a_typedval >> 24) & 0xFF
                a_data = struct.unpack_from('<I', data, ap + 12)[0]
                aname = pool.get(a_name)
                if a_rawval != 0xFFFFFFFF:
                    aval = pool.get(a_rawval)
                else:
                    aval = res_val(a_type, a_data, pool)
                attrs.append((aname, aval))
                ap += 20
            attr_str = ''.join(f' {k}="{v}"' for k, v in attrs)
            out.append('  ' * indent + f'<{name}{attr_str}>')
            indent += 1
        elif ctype == RES_XML_END_ELEMENT_TYPE:
            ns_idx, name_idx = struct.unpack_from('<II', data, off + 8 + 8)
            indent -= 1
            out.append('  ' * indent + f'</{pool.get(name_idx)}>')
        off += csize
        if csize == 0:
            break
    return '\n'.join(out)


if __name__ == '__main__':
    print(decode(sys.argv[1]))
