#!/usr/bin/env python3
"""
oeneyeVirtualLab.py - one front door for oeneye disk images.

  info   IMG                 inspect an image (boot signature, OENE magic,
                             partition table, filesystem, DOS vs Linux)
  run    IMG [--keys ...]    boot a DOS-type image (MOSFETQ DOS) inside the
                             built-in oeneyeVMx86 8086 emulator
  qemu   IMG [--serial]      boot a Linux-type image (or anything else) in
                             QEMU - the only way to run Linux, since
                             oeneyeVMx86 is an 8086 subset emulator
  ls     IMG                 list files in the image's FAT16 partition
  add    IMG FILE... [--out] copy files into the FAT16 partition (long names ok)
  build-linux [opts]         build oeneyeLinux.img by calling
                             build-oeneyeLinux.sh (needs internet + tools)

oeneyeVMx86.py must sit next to this file.
Author tag: @claude
"""
import argparse
import os
import shutil
import struct
import subprocess
import sys

HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, HERE)

MAGIC_OFFSET = 0x1E0  # "OENE" developer-edition magic in the boot sector
MAGIC = b"OENE"


# ---------------------------------------------------------------- inspect

def fat_kind(sector):
    """Guess FAT type from a boot sector's filesystem-type strings."""
    if sector[82:87] == b"FAT32":
        return "FAT32"
    for tag in (b"FAT12", b"FAT16"):
        if sector[54:62].startswith(tag):
            return tag.decode()
    return None


def inspect(data):
    info = {"size": len(data), "boot_signature": False, "oene_magic": False,
            "partitions": [], "fat": None, "syslinux": False, "kind": "unknown"}
    if len(data) < 512:
        return info
    info["boot_signature"] = data[510:512] == b"\x55\xaa"
    info["oene_magic"] = data[MAGIC_OFFSET:MAGIC_OFFSET + 4] == MAGIC

    # MBR partition table (only meaningful if the first sector isn't a
    # FAT boot sector itself)
    first_fat = fat_kind(data[:512])
    if not first_fat and info["boot_signature"]:
        for i in range(4):
            e = data[446 + 16 * i: 462 + 16 * i]
            status, ptype = e[0], e[4]
            lba, count = struct.unpack("<II", e[8:16])
            if ptype and count:
                info["partitions"].append(
                    {"n": i + 1, "active": status == 0x80,
                     "type": "0x%02X" % ptype, "start_lba": lba, "sectors": count})
    if first_fat:
        info["fat"] = first_fat
        info["syslinux"] = b"SYSLINUX" in data[:512]
    for p in info["partitions"]:
        off = p["start_lba"] * 512
        if off + 512 <= len(data):
            sec = data[off:off + 512]
            p["fat"] = fat_kind(sec)
            if b"SYSLINUX" in sec:
                info["syslinux"] = True

    if info["syslinux"] or b"GRUB" in data[:512]:
        info["kind"] = "linux"
    elif info["boot_signature"]:
        info["kind"] = "dos"   # includes MBR-style images whose boot code is MOSFETQ DOS
    return info


def cmd_info(args):
    data = open(args.image, "rb").read()
    i = inspect(data)
    print("image           :", args.image)
    print("size            : %d bytes (%.2f MiB)" % (i["size"], i["size"] / 1048576))
    print("boot signature  :", "55AA present" if i["boot_signature"] else "missing")
    print("OENE magic      :", "yes (0x1E0)" if i["oene_magic"] else "no")
    print("first-sector FAT:", i["fat"] or "-")
    for p in i["partitions"]:
        print("partition %d     : type %s%s start=%d sectors=%d fat=%s" % (
            p["n"], p["type"], " (active)" if p["active"] else "",
            p["start_lba"], p["sectors"], p.get("fat") or "-"))
    print("SYSLINUX        :", "yes" if i["syslinux"] else "no")
    print("detected as     :", i["kind"])
    return 0


# -------------------------------------------------------------------- run

def cmd_run(args):
    data = open(args.image, "rb").read()
    i = inspect(data)
    if i["kind"] == "linux":
        print("This looks like a Linux-type image. oeneyeVMx86 is an 8086 "
              "subset emulator and cannot run it.\nUse:  oeneyeVirtualLab.py "
              "qemu %s" % args.image, file=sys.stderr)
        return 2
    if not i["boot_signature"]:
        print("No 55AA boot signature - not a bootable image.", file=sys.stderr)
        return 2
    try:
        from oeneyeVMx86 import VM
    except ImportError:
        print("oeneyeVMx86.py not found next to oeneyeVirtualLab.py", file=sys.stderr)
        return 2
    vm = VM(data, trace=args.trace)
    vm.boot()
    for line in args.keys or []:
        vm.type_line(line)
    try:
        reason = vm.run(args.max_instr)
    except NotImplementedError as e:
        print(vm.screen_text())
        print("--- emulator stopped: %s" % e, file=sys.stderr)
        print("(the guest used an instruction outside the emulator's subset)",
              file=sys.stderr)
        return 1
    print(vm.screen_text())
    print("--- stopped: %s (%d instructions)" % (reason, vm.instr_count),
          file=sys.stderr)
    return 0


# ------------------------------------------------------------------- qemu

def cmd_qemu(args):
    q = shutil.which("qemu-system-x86_64") or shutil.which("qemu-system-i386")
    if not q:
        print("QEMU not found. Install it (e.g. apt install qemu-system-x86) "
              "and retry.", file=sys.stderr)
        return 2
    cmd = [q, "-m", str(args.mem),
           "-drive", "file=%s,format=raw,if=ide" % args.image]
    if args.serial:
        cmd += ["-nographic"]
    print("+", " ".join(cmd))
    return subprocess.call(cmd)


# ------------------------------------------------------------ build-linux

def cmd_build_linux(args):
    script = os.path.join(HERE, "build-oeneyeLinux.sh")
    if not os.path.exists(script):
        print("build-oeneyeLinux.sh not found next to this file.", file=sys.stderr)
        return 2
    env = dict(os.environ)
    if args.payload:
        env["PAYLOAD_DIR"] = args.payload
    if args.out:
        env["OUT_IMG"] = args.out
    if args.size:
        env["IMG_MB"] = str(args.size)
    return subprocess.call(["bash", script], env=env)



# ------------------------------------------------------------- FAT16 add

import time


def _lfn_checksum(short11):
    s = 0
    for c in short11:
        s = (((s & 1) << 7) + (s >> 1) + c) & 0xFF
    return s


def _short_name(name, taken):
    base, _, ext = name.rpartition(".") if "." in name else (name, "", "")
    clean = lambda t: "".join(c if (c.isalnum() or c in "_$~!#%&-{}()@'`^") else "_"
                              for c in t.upper().replace(" ", ""))
    base, ext = clean(base)[:8], clean(ext)[:3]
    cand = (base.ljust(8) + ext.ljust(3)).encode("ascii", "replace")
    if name == name.upper() and len(name.partition(".")[0]) <= 8 and cand not in taken \
            and len(name.rpartition(".")[2]) <= 3:
        return cand, False
    n = 1
    while True:
        tail = "~%d" % n
        cand = (base[:8 - len(tail)] + tail).ljust(8).encode() + ext.ljust(3).encode()
        if cand not in taken:
            return cand, True
        n += 1


class Fat16:
    """Minimal FAT16 file adder for an image's first FAT16 partition (or a
    superfloppy). Supports long file names; root directory only."""

    def __init__(self, data):
        self.d = data
        self.base = 0
        if fat_kind(data[:512]) is None:
            for i in range(4):
                e = data[446 + 16 * i: 462 + 16 * i]
                lba, cnt = struct.unpack("<II", e[8:16])
                if e[4] in (4, 6, 0xE, 0x0E) and cnt:
                    self.base = lba * 512
                    break
            else:
                raise ValueError("no FAT16 partition found")
        s = data[self.base:self.base + 512]
        self.bps = struct.unpack("<H", s[11:13])[0]
        self.spc = s[13]
        self.rsvd = struct.unpack("<H", s[14:16])[0]
        self.nfats = s[16]
        self.rootent = struct.unpack("<H", s[17:19])[0]
        tot = struct.unpack("<H", s[19:21])[0] or struct.unpack("<I", s[32:36])[0]
        self.spf = struct.unpack("<H", s[22:24])[0]
        if self.bps != 512:
            raise ValueError("only 512-byte sectors supported")
        self.root_sec = self.rsvd + self.nfats * self.spf
        self.data_sec = self.root_sec + (self.rootent * 32 + 511) // 512
        self.nclus = (tot - self.data_sec) // self.spc
        if not (4085 <= self.nclus < 65525):
            raise ValueError("not FAT16 (%d clusters)" % self.nclus)
        self.csize = self.spc * 512

    # raw helpers
    def _off(self, sec):
        return self.base + sec * 512

    def _fat_get(self, c):
        o = self._off(self.rsvd) + c * 2
        return struct.unpack("<H", self.d[o:o + 2])[0]

    def _fat_set(self, c, v):
        for f in range(self.nfats):
            o = self._off(self.rsvd + f * self.spf) + c * 2
            self.d[o:o + 2] = struct.pack("<H", v)

    def _root_entries(self):
        o = self._off(self.root_sec)
        return [(o + 32 * i, self.d[o + 32 * i: o + 32 * i + 32])
                for i in range(self.rootent)]

    def _free_chain(self, c):
        while 2 <= c < 0xFFF8:
            nxt = self._fat_get(c)
            self._fat_set(c, 0)
            c = nxt

    def listing(self):
        out, lfn = [], []
        for off, e in self._root_entries():
            if e[0] == 0:
                break
            if e[0] == 0xE5:
                lfn = []
                continue
            if e[11] == 0x0F:
                part = e[1:11] + e[14:26] + e[28:32]
                lfn.insert(0, part.decode("utf-16le").split("\0")[0].split("\uffff")[0])
                continue
            if e[11] & 8:
                lfn = []
                continue
            short = (e[:8].decode("ascii", "replace").strip() + "." +
                     e[8:11].decode("ascii", "replace").strip()).rstrip(".")
            out.append({"name": "".join(lfn) or short, "size": struct.unpack("<I", e[28:32])[0],
                        "cluster": struct.unpack("<H", e[26:28])[0]})
            lfn = []
        return out

    def read(self, name):
        for f in self.listing():
            if f["name"].lower() == name.lower():
                buf, c = bytearray(), f["cluster"]
                while 2 <= c < 0xFFF8:
                    o = self._off(self.data_sec + (c - 2) * self.spc)
                    buf += self.d[o:o + self.csize]
                    c = self._fat_get(c)
                return bytes(buf[:f["size"]])
        raise KeyError(name)

    def delete(self, name):
        lfn_offs = []
        for off, e in self._root_entries():
            if e[0] == 0:
                break
            if e[0] == 0xE5:
                lfn_offs = []
                continue
            if e[11] == 0x0F:
                lfn_offs.append(off)
                continue
            if not (e[11] & 8):
                full = None
                for f in self.listing():
                    if f["cluster"] == struct.unpack("<H", e[26:28])[0] and \
                            f["name"].lower() == name.lower():
                        full = f
                if full:
                    self._free_chain(full["cluster"])
                    for o in lfn_offs + [off]:
                        self.d[o] = 0xE5
                    return True
            lfn_offs = []
        return False

    def add(self, name, payload, replace=False):
        if any(f["name"].lower() == name.lower() for f in self.listing()):
            if not replace:
                raise FileExistsError(name)
            self.delete(name)
        need = max(1, -(-len(payload) // self.csize)) if payload else 0
        free = [c for c in range(2, self.nclus + 2) if self._fat_get(c) == 0]
        if len(free) < need:
            raise OSError("image full: need %d clusters, %d free" % (need, len(free)))
        chain = free[:need]
        for i, c in enumerate(chain):
            self._fat_set(c, chain[i + 1] if i + 1 < len(chain) else 0xFFFF)
            o = self._off(self.data_sec + (c - 2) * self.spc)
            piece = payload[i * self.csize:(i + 1) * self.csize]
            self.d[o:o + self.csize] = piece.ljust(self.csize, b"\0")
        taken = {bytes(e[:11]) for _, e in self._root_entries() if e[0] not in (0, 0xE5)}
        short, need_lfn = _short_name(name, taken)
        chunks = []
        if need_lfn:
            u = name.encode("utf-16le")
            u += b"\0\0" if len(name) % 13 else b""
            u = u.ljust(-(-len(u) // 26) * 26, b"\xff")
            chunks = [u[i:i + 26] for i in range(0, len(u), 26)]
        ents, ck = [], _lfn_checksum(short)
        for i in range(len(chunks) - 1, -1, -1):
            ch = chunks[i]
            seq = (i + 1) | (0x40 if i == len(chunks) - 1 else 0)
            ents.append(bytes([seq]) + ch[0:10] + b"\x0f\x00" + bytes([ck]) +
                        ch[10:22] + b"\x00\x00" + ch[22:26])
        t = time.localtime()
        fdate = ((t.tm_year - 1980) << 9) | (t.tm_mon << 5) | t.tm_mday
        ftime = (t.tm_hour << 11) | (t.tm_min << 5) | (t.tm_sec // 2)
        ents.append(short + b"\x20" + b"\0" * 10 +
                    struct.pack("<HHHI", ftime, fdate, chain[0] if chain else 0, len(payload)))
        # find contiguous free slots (0x00 or 0xE5) in the root directory
        slots = self._root_entries()
        run = 0
        for i, (off, e) in enumerate(slots):
            run = run + 1 if e[0] in (0, 0xE5) else 0
            if run == len(ents):
                start = i - len(ents) + 1
                for k, ent in enumerate(ents):
                    o = slots[start + k][0]
                    self.d[o:o + 32] = ent
                return
        raise OSError("root directory full")


def cmd_add(args):
    data = bytearray(open(args.image, "rb").read())
    fs = Fat16(data)
    for path in args.files:
        name = os.path.basename(path)
        fs.add(name, open(path, "rb").read(), replace=args.replace)
        print("added", name)
    out = args.out or args.image
    open(out, "wb").write(bytes(data))
    print("wrote", out)
    return 0


def cmd_ls(args):
    fs = Fat16(bytearray(open(args.image, "rb").read()))
    for f in fs.listing():
        print("%10d  %s" % (f["size"], f["name"]))
    return 0


# ------------------------------------------------------------------- main

def main(argv=None):
    ap = argparse.ArgumentParser(prog="oeneyeVirtualLab.py",
                                 description="oeneye image lab")
    sub = ap.add_subparsers(dest="cmd", required=True)

    p = sub.add_parser("info", help="inspect an image")
    p.add_argument("image")
    p.set_defaults(fn=cmd_info)

    p = sub.add_parser("run", help="boot a DOS image in oeneyeVMx86")
    p.add_argument("image")
    p.add_argument("--keys", nargs="*", metavar="LINE",
                   help='lines to type, e.g. --keys HELP VER')
    p.add_argument("--max-instr", type=int, default=None)
    p.add_argument("--trace", action="store_true")
    p.set_defaults(fn=cmd_run)

    p = sub.add_parser("qemu", help="boot an image in QEMU")
    p.add_argument("image")
    p.add_argument("--mem", type=int, default=256, help="RAM in MB")
    p.add_argument("--serial", action="store_true",
                   help="text-only, console on this terminal")
    p.set_defaults(fn=cmd_qemu)

    p = sub.add_parser("ls", help="list files in the image's FAT16 partition")
    p.add_argument("image")
    p.set_defaults(fn=cmd_ls)

    p = sub.add_parser("add", help="copy files into the image's FAT16 partition")
    p.add_argument("image")
    p.add_argument("files", nargs="+")
    p.add_argument("--out", help="write to a new image instead of in place")
    p.add_argument("--replace", action="store_true", help="overwrite same-named files")
    p.set_defaults(fn=cmd_add)

    p = sub.add_parser("build-linux", help="build oeneyeLinux.img")
    p.add_argument("--payload", help="dir of oeneye files to include")
    p.add_argument("--out", help="output image (default oeneyeLinux.img)")
    p.add_argument("--size", type=int, help="image size in MB (default 64)")
    p.set_defaults(fn=cmd_build_linux)

    args = ap.parse_args(argv)
    return args.fn(args)


if __name__ == "__main__":
    sys.exit(main())
