#!/usr/bin/env python
"""Quartus BitRock/CookFS Extractor — passive extraction without execution.

Extracts files from Intel Quartus .run installers by parsing the CookFS
virtual filesystem embedded in the BitRock installer binary.  Remaps internal
BitRock component paths to the final Quartus installation layout using the
embedded project.xml folder definitions.

Usage:
    quartus-extractor.py <installer.run> --list
    quartus-extractor.py <installer.run> --list --component quartus_base
    quartus-extractor.py <installer.run> --extract <outdir>
    quartus-extractor.py <installer.run> --extract <outdir> --component quartus_base
    quartus-extractor.py <installer.run> --extract <outdir> --raw
"""

from __future__ import annotations

import argparse
import lzma
import os
import re
import stat
import struct
import sys
import zlib
from typing import BinaryIO

# -- Types ------------------------------------------------------------------

# (kind, path, mtime, blocks)  where blocks = [(page, offset, size), ...]
Entry = tuple[str, str, int, list[tuple[int, int, int]]]

# -- Constants --------------------------------------------------------------

CFS_MAGIC = b"CFS0002"
CFS_SUFFIX_SIZE = 16
CFS_INDEX_MAGIC = b"CFS2.200"


# -- Folder mapping ---------------------------------------------------------

def find_folder_mapping(fobj: BinaryIO, cfs_endoffset: int) -> dict[str, str]:
    """Extract folder-to-destination mapping from project.xml."""
    fobj.seek(cfs_endoffset)
    data: bytes = fobj.read(10 * 1024 * 1024)

    mapping: dict[str, str] = {}
    for i in range(len(data) - 2):
        if data[i] != 0x78:
            continue
        if ((data[i] << 8) + data[i + 1]) % 31 != 0:
            continue
        try:
            dec = zlib.decompress(data[i : i + 1024 * 1024])
        except zlib.error:
            continue
        if b"<project>" not in dec or b"<folder>" not in dec:
            continue
        for folder in re.findall(b"<folder>.*?</folder>", dec, re.DOTALL):
            name_m = re.search(b"<name>(.*?)</name>", folder)
            dest_m = re.search(b"<destination>(.*?)</destination>", folder)
            if name_m and dest_m:
                name = name_m.group(1).decode()
                dest = dest_m.group(1).decode()
                dest = dest.replace("${installdir}", ".")
                dest = dest.replace("${product_key}", "quartus")
                dest = dest.strip("/").lstrip(".").strip("/")
                mapping[name] = dest
        break

    return mapping


def remap_path(vfs_path: str, folder_mapping: dict[str, str]) -> str:
    """Remap a CookFS VFS path to the final installation path."""
    parts = vfs_path.split("/")
    if len(parts) < 3:
        return vfs_path

    folder = parts[1]
    subpath = "/".join(parts[2:])

    if folder in folder_mapping:
        dest = folder_mapping[folder]
        return f"{dest}/{subpath}" if dest else subpath

    return "/".join(parts[1:])


# -- CookFS parser ----------------------------------------------------------

def find_cfs_suffix(fobj: BinaryIO) -> tuple[int, int, int, int]:
    """Find the CFS0002 suffix by searching backward from EOF."""
    fobj.seek(0, 2)
    file_size: int = fobj.tell()

    search_size = min(file_size, 20 * 1024 * 1024)
    fobj.seek(file_size - search_size)
    data: bytes = fobj.read(search_size)

    idx = data.rfind(CFS_MAGIC)
    if idx < 0:
        raise ValueError("CFS0002 magic not found")

    suffix_abs = (file_size - search_size) + idx - 9
    fobj.seek(suffix_abs)
    suffix: bytes = fobj.read(CFS_SUFFIX_SIZE)

    index_size, num_pages = struct.unpack(">II", suffix[:8])
    compression_id: int = suffix[8]
    magic = suffix[9:]
    if magic != CFS_MAGIC:
        raise ValueError(f"CFS suffix magic mismatch: {magic!r}")

    endoffset = suffix_abs + CFS_SUFFIX_SIZE
    return index_size, num_pages, compression_id, endoffset


def parse_layout(
    fobj: BinaryIO, index_size: int, num_pages: int, endoffset: int,
) -> tuple[list[int], list[int], int]:
    """Calculate CookFS page layout from suffix parameters."""
    md5_offset = endoffset - 16 - index_size - (num_pages * 4) - (num_pages * 16)
    sizes_offset = md5_offset + (num_pages * 16)
    blob_offset = endoffset - 16 - index_size

    fobj.seek(sizes_offset)
    page_sizes: list[int] = [
        struct.unpack_from(">I", fobj.read(4))[0] for _ in range(num_pages)
    ]

    page_offsets: list[int] = []
    off = md5_offset - sum(page_sizes)
    for psize in page_sizes:
        page_offsets.append(off)
        off += psize

    return page_offsets, page_sizes, blob_offset


def parse_index(
    fobj: BinaryIO, blob_offset: int, index_size: int,
) -> list[Entry]:
    """Read and decompress the CFS2.200 filesystem index."""
    fobj.seek(blob_offset)
    blob: bytes = fobj.read(index_size + 1)
    comp_id: int = blob[0]
    payload: bytes = blob[1:]

    decoded: bytes
    if comp_id == 0:
        decoded = payload
    elif comp_id == 1:
        decoded = zlib.decompress(payload, wbits=-15)
    elif comp_id == 255:
        decoded = lzma.decompress(payload, format=lzma.FORMAT_ALONE)
    else:
        raise ValueError(f"Unknown index compression: {comp_id}")

    if not decoded.startswith(CFS_INDEX_MAGIC):
        raise ValueError(f"Invalid index magic: {decoded[:8]!r}")

    pos = 8
    entries: list[Entry] = []

    def read_u32() -> int:
        nonlocal pos
        val = struct.unpack_from(">I", decoded, pos)[0]
        pos += 4
        return int(val)

    def read_i32() -> int:
        nonlocal pos
        val = struct.unpack_from(">i", decoded, pos)[0]
        pos += 4
        return int(val)

    def read_u64() -> int:
        nonlocal pos
        val = struct.unpack_from(">Q", decoded, pos)[0]
        pos += 8
        return int(val)

    def read_dir(prefix: str) -> None:
        nonlocal pos
        if pos + 4 > len(decoded):
            return
        child_count = read_u32()
        for _ in range(child_count):
            name_len = decoded[pos]
            pos += 1
            name = decoded[pos : pos + name_len].decode("utf-8", "replace")
            pos += name_len + 1
            mtime = read_u64()
            block_count = read_i32()
            full = f"{prefix}/{name}" if prefix else name
            if block_count == -1:
                entries.append(("dir", full, mtime, []))
                read_dir(full)
            else:
                blocks: list[tuple[int, int, int]] = [
                    (read_u32(), read_u32(), read_u32())
                    for _ in range(block_count)
                ]
                entries.append(("file", full, mtime, blocks))

    read_dir("")
    return entries


# -- Page decompression -----------------------------------------------------

def decompress_page(raw: bytes) -> bytes:
    """Decompress a CookFS page based on its tag byte."""
    if not raw:
        return b""
    tag: int = raw[0]
    payload: bytes = raw[1:]
    if tag == 0:
        return payload
    if tag == 1:
        return zlib.decompress(payload, wbits=-15)
    if tag == 255:
        return lzma.decompress(payload, format=lzma.FORMAT_ALONE)
    raise ValueError(f"Unknown page compression tag: {tag}")


# -- Extraction helpers -----------------------------------------------------

def extract_file_data(
    fobj: BinaryIO,
    entry: Entry,
    page_offsets: list[int],
    page_sizes: list[int],
    page_cache: dict[int, bytes],
) -> bytes:
    """Extract a single file's content from CookFS pages."""
    data = b""
    for page_idx, offset_in_page, size in entry[3]:
        if page_idx not in page_cache:
            fobj.seek(page_offsets[page_idx])
            raw = fobj.read(page_sizes[page_idx])
            page_cache[page_idx] = decompress_page(raw)
        page = page_cache[page_idx]
        data += page[offset_in_page : offset_in_page + size]
    return data


def set_exec_permission(filepath: str, data: bytes) -> None:
    """Set executable permission based on file magic bytes."""
    basename = os.path.basename(filepath)
    if data[:4] == b"\x7fELF" or data[:2] == b"#!":
        os.chmod(filepath, 0o755)
    elif basename.endswith(".so") or ".so." in basename:
        os.chmod(filepath, 0o755)


def fix_permissions(target_dir: str) -> int:
    """Set +x on ELF binaries and shebang scripts in a directory."""
    count = 0
    for root, _dirs, files in os.walk(target_dir):
        for fname in files:
            filepath = os.path.join(root, fname)
            try:
                with open(filepath, "rb") as fh:
                    magic = fh.read(4)
                if magic[:4] == b"\x7fELF" or magic[:2] == b"#!":
                    os.chmod(
                        filepath,
                        os.stat(filepath).st_mode
                        | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH,
                    )
                    count += 1
                elif fname.endswith(".so") or ".so." in fname:
                    os.chmod(
                        filepath,
                        os.stat(filepath).st_mode
                        | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH,
                    )
                    count += 1
            except OSError:
                pass
    return count


def collect_bigfile_blocks(
    files: list[Entry], base_path: str, base_blocks: list[tuple[int, int, int]],
) -> list[tuple[int, int, int]]:
    """Collect BigFile chunk blocks and append to base blocks."""
    all_blocks = list(base_blocks)
    chunk_idx = 1
    while True:
        chunk_name = f"{base_path}___bitrockBigFile{chunk_idx}"
        chunk_entry = next((e for e in files if e[1] == chunk_name), None)
        if chunk_entry is None:
            break
        all_blocks.extend(chunk_entry[3])
        chunk_idx += 1
    return all_blocks


def do_extract(
    fobj: BinaryIO,
    files: list[Entry],
    folder_mapping: dict[str, str],
    page_offsets: list[int],
    page_sizes: list[int],
    args: argparse.Namespace,
) -> None:
    """Extract files from the CookFS archive."""
    os.makedirs(args.extract, exist_ok=True)
    page_cache: dict[int, bytes] = {}
    total_files = len(files)
    extracted = 0

    for i, entry in enumerate(files):
        _, path, mtime, blocks = entry

        if "___bitrockBigFile" in path:
            continue

        all_blocks = collect_bigfile_blocks(files, path, blocks)
        mapped = remap_path(path, folder_mapping) if folder_mapping else path
        out_path = os.path.join(args.extract, mapped)
        os.makedirs(os.path.dirname(out_path), exist_ok=True)

        stitched: Entry = ("file", path, mtime, all_blocks)
        data = extract_file_data(
            fobj, stitched, page_offsets, page_sizes, page_cache,
        )

        with open(out_path, "wb") as out:
            out.write(data)

        set_exec_permission(out_path, data)
        extracted += 1

        if (i + 1) % 1000 == 0 or i + 1 == total_files:
            print(f"  [{i + 1}/{total_files}] {mapped}", file=sys.stderr)

        if args.limit and extracted >= args.limit:
            break

    print(f"\nExtracted {extracted} files to {args.extract}", file=sys.stderr)


# -- Main -------------------------------------------------------------------

def main() -> None:
    """Main entry point."""
    parser = argparse.ArgumentParser(
        description="Extract Quartus .run installer (passive, no execution)",
    )
    parser.add_argument("installer", help="Path to .run file or directory")
    parser.add_argument("--list", action="store_true", help="List files")
    parser.add_argument("--extract", metavar="DIR", help="Extract to directory")
    parser.add_argument("--component", help="Filter by component")
    parser.add_argument("--limit", type=int, help="Limit output entries")
    parser.add_argument("--stats", action="store_true", help="Show statistics")
    parser.add_argument("--raw", action="store_true", help="Skip path remapping")
    parser.add_argument("--fix-permissions", action="store_true", help="Fix +x bits")
    args = parser.parse_args()

    if args.fix_permissions:
        count = fix_permissions(args.installer)
        print(f"Fixed permissions on {count} files", file=sys.stderr)
        return

    if not args.list and not args.extract and not args.stats:
        parser.error("Specify --list, --extract, --stats, or --fix-permissions")

    with open(args.installer, "rb") as fobj:
        index_size, num_pages, _comp_id, endoffset = find_cfs_suffix(fobj)
        print(f"CookFS: {num_pages} pages, index {index_size} bytes", file=sys.stderr)

        page_offsets, page_sizes, blob_offset = parse_layout(
            fobj, index_size, num_pages, endoffset,
        )
        entries = parse_index(fobj, blob_offset, index_size)

        folder_mapping: dict[str, str] = {}
        if not args.raw:
            folder_mapping = find_folder_mapping(fobj, endoffset)
            if folder_mapping:
                print(
                    f"Path mapping: {len(folder_mapping)} folders", file=sys.stderr,
                )
            else:
                print("WARNING: No path mapping found", file=sys.stderr)

        dirs = [e for e in entries if e[0] == "dir"]
        files = [e for e in entries if e[0] == "file"]
        print(f"Index: {len(dirs)} dirs, {len(files)} files", file=sys.stderr)

        if args.component:
            entries = [e for e in entries if e[1].startswith(args.component)]
            files = [e for e in entries if e[0] == "file"]
            dirs = [e for e in entries if e[0] == "dir"]
            print(
                f"Filtered to {args.component}:"
                f" {len(dirs)} dirs, {len(files)} files",
                file=sys.stderr,
            )

        if args.stats:
            components: dict[str, dict[str, int]] = {}
            for entry in entries:
                comp = entry[1].split("/")[0]
                if comp not in components:
                    components[comp] = {"dirs": 0, "files": 0, "bytes": 0}
                if entry[0] == "dir":
                    components[comp]["dirs"] += 1
                else:
                    components[comp]["bytes"] += sum(b[2] for b in entry[3])
                    components[comp]["files"] += 1
            print(f"\n{'Component':<25} {'Files':>8} {'Dirs':>6} {'Size':>12}")
            print("-" * 55)
            for comp in sorted(components):
                info = components[comp]
                size_mb = info["bytes"] / 1024 / 1024
                print(
                    f"{comp:<25} {info['files']:>8}"
                    f" {info['dirs']:>6} {size_mb:>9.1f} MB",
                )

        if args.list:
            limit = args.limit or len(entries)
            for entry in entries[:limit]:
                kind, path, _mtime, blocks = entry
                mapped = remap_path(path, folder_mapping) if folder_mapping else path
                if kind == "file":
                    total = sum(b[2] for b in blocks)
                    print(f"{mapped}\t{total}")
                else:
                    print(f"{mapped}/")

        if args.extract:
            do_extract(fobj, files, folder_mapping, page_offsets, page_sizes, args)

if __name__ == "__main__":
    main()
