#!/usr/bin/env python3
"""Safely and recursively unpack one feedback payload on the target host."""

import argparse
import json
import os
import shutil
import stat
import tarfile
import zipfile


SUPPORTED_LOG_EXTENSIONS = (".qdss", ".qmdl", ".hdf", ".isf")
ARCHIVE_SUFFIXES = (
    ".zip", ".tar", ".tar.gz", ".tgz", ".tar.bz2", ".tbz2",
    ".tar.xz", ".txz", ".gz", ".bz2", ".xz", ".7z", ".rar",
)


class ExtractError(Exception):
    pass


class Budget:
    def __init__(self, max_files, max_bytes, reserve_bytes):
        self.max_files = max_files
        self.max_bytes = max_bytes
        self.reserve_bytes = reserve_bytes
        self.files = 0
        self.bytes = 0
        self.archives = 0

    def add(self, size, destination):
        self.files += 1
        self.bytes += max(0, int(size))
        if self.files > self.max_files or self.bytes > self.max_bytes:
            raise ExtractError("feedback archive exceeds extraction limits")
        if shutil.disk_usage(destination).free < size + self.reserve_bytes:
            raise ExtractError("insufficient target disk space")


def safe_destination(root, name):
    normalized = name.replace("\\", "/")
    if normalized.startswith("/"):
        raise ExtractError("absolute archive path is not allowed")
    parts = [part for part in normalized.split("/") if part not in ("", ".")]
    if not parts or ".." in parts:
        raise ExtractError("archive path traversal is not allowed")
    destination = os.path.abspath(os.path.join(root, *parts))
    root = os.path.abspath(root)
    if os.path.commonpath([root, destination]) != root:
        raise ExtractError("archive path escapes destination")
    return destination


def archive_kind(pathname):
    if os.path.getsize(pathname) == 0:
        return None
    with open(pathname, "rb") as fh:
        magic = fh.read(8)
    if zipfile.is_zipfile(pathname):
        return "zip"
    if tarfile.is_tarfile(pathname):
        return "tar"
    lower = pathname.lower()
    if magic.startswith((b"BZh", b"\x1f\x8b", b"\xfd7zXZ\x00")):
        raise ExtractError("single-file compressed archive is not supported")
    if lower.endswith((".7z", ".rar")):
        raise ExtractError("unsupported nested archive")
    return None


def extract_zip(source, destination, budget):
    seen = set()
    with zipfile.ZipFile(source) as archive:
        for member in archive.infolist():
            target = safe_destination(destination, member.filename)
            if target in seen:
                raise ExtractError("duplicate archive member")
            seen.add(target)
            if member.flag_bits & 0x1:
                raise ExtractError("encrypted archive is not supported")
            mode = (member.external_attr >> 16) & 0xFFFF
            if mode and stat.S_ISLNK(mode):
                raise ExtractError("archive links are not allowed")
            file_type = stat.S_IFMT(mode) if mode else 0
            if file_type and file_type not in (stat.S_IFREG, stat.S_IFDIR):
                raise ExtractError("archive special files are not allowed")
            if member.is_dir():
                os.makedirs(target, exist_ok=True)
                continue
            budget.add(member.file_size, destination)
            os.makedirs(os.path.dirname(target), exist_ok=True)
            with archive.open(member) as src, open(target, "wb") as dst:
                shutil.copyfileobj(src, dst, 4 * 1024 * 1024)


def extract_tar(source, destination, budget):
    seen = set()
    with tarfile.open(source, "r:*") as archive:
        for member in archive:
            target = safe_destination(destination, member.name)
            if target in seen:
                raise ExtractError("duplicate archive member")
            seen.add(target)
            if member.issym() or member.islnk() or member.isdev():
                raise ExtractError("archive links and devices are not allowed")
            if member.isdir():
                os.makedirs(target, exist_ok=True)
                continue
            if not member.isfile():
                continue
            budget.add(member.size, destination)
            os.makedirs(os.path.dirname(target), exist_ok=True)
            src = archive.extractfile(member)
            if src is None:
                raise ExtractError("failed to read archive member")
            with src, open(target, "wb") as dst:
                shutil.copyfileobj(src, dst, 4 * 1024 * 1024)


def extract_one(source, destination, kind, budget):
    os.makedirs(destination, exist_ok=True)
    budget.archives += 1
    if kind == "zip":
        extract_zip(source, destination, budget)
    else:
        extract_tar(source, destination, budget)


def remove_stale_outputs(root):
    for current, directories, _files in os.walk(root, topdown=True):
        for name in list(directories):
            if name == "auto_analysis":
                path = os.path.join(current, name)
                shutil.rmtree(path)
                directories.remove(name)


def recursive_unpack(source, destination, max_depth, budget):
    outer_kind = archive_kind(source)
    if outer_kind is None:
        budget.add(os.path.getsize(source), destination)
        os.makedirs(destination, exist_ok=True)
        shutil.copy2(source, os.path.join(destination, os.path.basename(source)))
        return
    extract_one(source, destination, outer_kind, budget)
    remove_stale_outputs(destination)
    queue = [(destination, 1)]
    while queue:
        scan_root, depth = queue.pop(0)
        nested = []
        for current, directories, files in os.walk(scan_root):
            directories[:] = sorted(
                directory for directory in directories
                if directory != "auto_analysis")
            for name in sorted(files):
                path = os.path.join(current, name)
                kind = archive_kind(path)
                if kind:
                    nested.append((path, kind))
        for path, kind in nested:
            if depth >= max_depth:
                raise ExtractError("nested archive depth exceeds limit")
            unpacked = path + ".unpacked"
            if os.path.exists(unpacked):
                raise ExtractError("nested extraction destination already exists")
            extract_one(path, unpacked, kind, budget)
            remove_stale_outputs(unpacked)
            os.remove(path)
            queue.append((unpacked, depth + 1))


def find_analysis_input(root):
    log_directories = []
    for current, directories, files in os.walk(root):
        directories[:] = sorted(d for d in directories if d != "auto_analysis")
        if any(os.path.splitext(name)[1].lower() in SUPPORTED_LOG_EXTENSIONS
               for name in files):
            log_directories.append(current)
    if not log_directories:
        raise ExtractError("no supported QCAT log files found")
    if len(log_directories) > 1:
        raise ExtractError("multiple independent QCAT log directories found")
    return log_directories[0]


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--source", required=True)
    parser.add_argument("--destination", required=True)
    parser.add_argument("--max-depth", type=int, default=4)
    parser.add_argument("--max-files", type=int, required=True)
    parser.add_argument("--max-bytes", type=int, required=True)
    parser.add_argument("--reserve-bytes", type=int, required=True)
    args = parser.parse_args()

    if args.max_depth < 1 or args.max_depth > 10:
        raise ExtractError("invalid extraction depth")
    if os.path.exists(args.destination):
        shutil.rmtree(args.destination)
    budget = Budget(args.max_files, args.max_bytes, args.reserve_bytes)
    recursive_unpack(args.source, args.destination, args.max_depth, budget)
    remove_stale_outputs(args.destination)
    analysis_input = find_analysis_input(args.destination)
    print(json.dumps({
        "analysis_input": analysis_input,
        "archives_extracted": budget.archives,
        "files_extracted": budget.files,
        "bytes_extracted": budget.bytes,
    }, sort_keys=True))


if __name__ == "__main__":
    try:
        main()
    except ExtractError as exc:
        raise SystemExit("ERROR: {}".format(exc))
