#!/usr/bin/env python3
# trdos 1.0: a Midnight Commander extfs helper for TR-DOS disk images (.trd) and SCL archives (.scl).
# Copyright (c) 2026 Spectre (Optical Brothers), https://www.zxby.org. MIT License (see LICENSE).
#
# mc runs it as:  trdos list ARCHIVE
#                 trdos copyout ARCHIVE NAME FILE
#                 trdos copyin ARCHIVE NAME FILE
#                 trdos rm ARCHIVE NAME
# and for F3 on an image (mc.ext.ini's View):  trdos catalog ARCHIVE
#
# The archive's root holds the files as they are on the disk ("raw"), hobeta/ the same files with a hobeta header,
# CATALOG the disk's catalog as text. See README.md.

import os
import re
import stat
import struct
import sys
import tempfile
import time

SECTOR = 256
MAX_FILES = 128
MAX_SECTORS = 255  # a TR-DOS file holds 255 sectors (65280 bytes) at most
HOBETA_DIR = 'hobeta'
CATALOG = 'CATALOG'
HOBETA_HEADER = 17

# The disk types of sector 8 (#8E3) and their sizes in sectors.
TRD_TYPES = {
    0x16: (2560, '80 tracks, 2 sides'),
    0x17: (1280, '40 tracks, 2 sides'),
    0x18: (1280, '80 tracks, 1 side'),
    0x19: (640, '40 tracks, 1 side'),
}
TRD_SIZE = 2560 * SECTOR

# Sector 8: the disk's info, at these offsets of the image.
FIRST_FREE_SECTOR = 0x8E1
FIRST_FREE_TRACK = 0x8E2
DISK_TYPE = 0x8E3
FILE_COUNT = 0x8E4  # the files in the catalog, the deleted ones too
FREE_SECTORS = 0x8E5
TRDOS_ID = 0x8E7  # #10 on a TR-DOS disk
DELETED_COUNT = 0x8F4
LABEL = 0x8F5

# The three-character extensions ("song.pt3"): the type and the two bytes of the start are these characters. QC's
# packed files keep the original type after "_" ("game.p_C": Hrust, "game.z_C": ZX0), so "_" and a printable
# character are taken too.
EXT3 = b'abcdefghijklmnopqrstuvwxyz0123456789'


class Error(Exception):
    pass


class Entry(object):
    """A file of a catalog. name is the 8 bytes as stored; start is the TRD's first logical sector."""

    def __init__(self, name, type_, p1, p2, sectors, start=0, index=0, deleted=False):
        self.name = name
        self.type = type_
        self.p1 = p1  # the start (BASIC: the length with the variables)
        self.p2 = p2  # the length (BASIC: the program's length)
        self.sectors = sectors
        self.start = start
        self.index = index
        self.deleted = deleted

    def header(self):
        """The 13 bytes TRD, SCL and hobeta have alike: the name, the type, the start and the length."""
        return self.name + bytes([self.type]) + struct.pack('<HH', self.p1, self.p2)

    def same(self, other):
        return self.header() == other.header() and self.sectors == other.sectors


def u16(data, offset):
    return struct.unpack('<H', bytes(data[offset:offset + 2]))[0]


class TRD(object):
    kind = 'TRD'

    def __init__(self, data):
        if len(data) < 9 * SECTOR or data[TRDOS_ID] != 0x10:
            raise Error('not a TR-DOS disk image (no #10 at #8E7)')
        self.data = bytearray(data)
        self.read_catalog()

    @classmethod
    def blank(cls):
        """A disk as TR-DOS formats it: 80 tracks, 2 sides, nothing on it."""
        data = bytearray(TRD_SIZE)
        data[FIRST_FREE_TRACK] = 1
        data[DISK_TYPE] = 0x16
        data[FREE_SECTORS:FREE_SECTORS + 2] = struct.pack('<H', 2560 - 16)
        data[TRDOS_ID] = 0x10
        data[0x8EA:0x8F3] = b' ' * 9
        data[LABEL:LABEL + 8] = b' ' * 8
        return cls(data)

    def read_catalog(self):
        self.entries = []
        for i in range(MAX_FILES):
            raw = self.data[i * 16:i * 16 + 16]
            if raw[0] == 0:
                break
            p1, p2 = struct.unpack('<HH', bytes(raw[9:13]))
            self.entries.append(Entry(bytes(raw[0:8]), raw[8], p1, p2, raw[13], start=raw[15] * 16 + raw[14],
                                      index=i, deleted=raw[0] == 1))

    def files(self):
        return [e for e in self.entries if not e.deleted]

    @property
    def first_free(self):
        return self.data[FIRST_FREE_TRACK] * 16 + self.data[FIRST_FREE_SECTOR]

    @first_free.setter
    def first_free(self, n):
        self.data[FIRST_FREE_SECTOR], self.data[FIRST_FREE_TRACK] = n % 16, n // 16

    @property
    def free(self):
        return u16(self.data, FREE_SECTORS)

    @free.setter
    def free(self, n):
        self.data[FREE_SECTORS:FREE_SECTORS + 2] = struct.pack('<H', max(0, min(n, 0xFFFF)))

    @property
    def total(self):
        """The disk's size in sectors: by its type, or, for a type unknown, as much as sector 8 counts."""
        if self.data[DISK_TYPE] in TRD_TYPES:
            return TRD_TYPES[self.data[DISK_TYPE]][0]
        return min(self.first_free + self.free, 256 * 16)

    def body(self, e):
        """The file's sectors; past the end of a cut image, zeros."""
        b = bytes(self.data[e.start * SECTOR:(e.start + e.sectors) * SECTOR])
        return b + bytes(e.sectors * SECTOR - len(b))

    def put(self, start, body):
        end = start * SECTOR + len(body)
        if len(self.data) < end:
            self.data.extend(bytes(end - len(self.data)))
        self.data[start * SECTOR:end] = body

    def add(self, e, body):
        """Writes the file after the last one, as TR-DOS does: the files already there never move."""
        n = len(body) // SECTOR
        if len(self.entries) >= MAX_FILES:
            raise Error('the catalog is full: %d files' % MAX_FILES)
        first = self.first_free
        start = max([first] + [x.start + x.sectors for x in self.entries])
        need = start + n - first  # what the free sectors lose
        if need > self.free or start + n > self.total:
            raise Error('no room on the disk: %d sectors needed, %d free' % (n, max(0, self.free - (start - first))))
        i = len(self.entries)
        self.data[i * 16:i * 16 + 16] = e.header() + bytes([n, start % 16, start // 16])
        if i + 1 < MAX_FILES:  # the catalog ends here, as after MOVE: old entries past it stay gone
            self.data[(i + 1) * 16] = 0
        self.put(start, body)
        self.first_free = start + n
        self.free = self.free - need
        self.data[FILE_COUNT] = i + 1
        self.read_catalog()

    def delete(self, e):
        """Deletes the file as TR-DOS does: marks it; only the deleted files at the catalog's end give their room
        back (when nothing lies between them and the free sectors)."""
        self.data[e.index * 16] = 1
        self.data[DELETED_COUNT] = (self.data[DELETED_COUNT] + 1) & 0xFF
        self.read_catalog()
        while self.entries and self.entries[-1].deleted:
            last = self.entries[-1]
            if last.start + last.sectors != self.first_free:
                break
            self.data[last.index * 16:last.index * 16 + 16] = bytes(16)
            self.first_free = last.start
            self.free = self.free + last.sectors
            self.data[FILE_COUNT] = last.index
            self.data[DELETED_COUNT] = max(0, self.data[DELETED_COUNT] - 1)
            self.entries.pop()

    def replace(self, e, new, body):
        """Overwrites a file: in place when the sectors are as many, else the old one deleted and the new added."""
        if len(body) // SECTOR == e.sectors:
            self.data[e.index * 16:e.index * 16 + 13] = new.header()
            self.put(e.start, body)
            self.read_catalog()
        else:
            self.delete(e)
            self.add(new, body)

    def bytes(self):
        return bytes(self.data)

    def describe(self):
        t = self.data[DISK_TYPE]
        lines = ['TR-DOS disk, label "%s", %s (#%02X)' % (show_text(self.data[LABEL:LABEL + 8]),
                                                          TRD_TYPES.get(t, (0, 'type unknown'))[1], t)]
        lines.append('%d files, %d deleted; %d sectors free from track %d, sector %d' % (
            len(self.entries), sum(1 for e in self.entries if e.deleted), self.free,
            self.data[FIRST_FREE_TRACK], self.data[FIRST_FREE_SECTOR]))
        return lines


class SCL(object):
    kind = 'SCL'

    def __init__(self, data):
        if len(data) < 9 or data[:8] != b'SINCLAIR':
            raise Error('not an SCL archive')
        n = data[8]
        pos = 9 + 14 * n
        if pos > len(data):
            raise Error('the catalog of %d files is past the end of the SCL' % n)
        self.entries, self.bodies = [], []
        for i in range(n):
            h = data[9 + 14 * i:9 + 14 * (i + 1)]
            p1, p2 = struct.unpack('<HH', bytes(h[9:13]))
            e = Entry(bytes(h[0:8]), h[8], p1, p2, h[13], index=i)
            body = bytes(data[pos:pos + e.sectors * SECTOR])
            self.entries.append(e)
            self.bodies.append(body + bytes(e.sectors * SECTOR - len(body)))  # a cut SCL: zeros
            pos += e.sectors * SECTOR
        self.checksum_ok = len(data) == pos + 4 and u16(data, pos) + (u16(data, pos + 2) << 16) == sum(data[:pos]) & 0xFFFFFFFF

    @classmethod
    def blank(cls):
        return cls(scl_bytes([], []))

    def files(self):
        return list(self.entries)

    def body(self, e):
        return self.bodies[e.index]

    def reindex(self):
        for i, e in enumerate(self.entries):
            e.index = i

    def add(self, e, body):
        if len(self.entries) >= MAX_FILES:
            raise Error('the catalog is full: %d files' % MAX_FILES)
        self.entries.append(Entry(e.name, e.type, e.p1, e.p2, len(body) // SECTOR))
        self.bodies.append(body)
        self.reindex()

    def delete(self, e):
        del self.entries[e.index], self.bodies[e.index]
        self.reindex()

    def replace(self, e, new, body):
        self.entries[e.index] = Entry(new.name, new.type, new.p1, new.p2, len(body) // SECTOR)
        self.bodies[e.index] = body
        self.reindex()

    def bytes(self):
        return scl_bytes(self.entries, self.bodies)

    def describe(self):
        return ['SCL archive, %d files, %d sectors; checksum %s' % (
            len(self.entries), sum(e.sectors for e in self.entries), 'right' if self.checksum_ok else 'WRONG')]


def scl_bytes(entries, bodies):
    out = bytearray(b'SINCLAIR')
    out.append(len(entries))
    for e in entries:
        out += e.header() + bytes([e.sectors])
    for b in bodies:
        out += b
    out += struct.pack('<I', sum(out) & 0xFFFFFFFF)
    return bytes(out)


def scl_size_ok(data):
    n = data[8]
    if len(data) < 9 + 14 * n:
        return False
    return len(data) == 9 + 14 * n + sum(data[9 + 14 * i + 13] for i in range(n)) * SECTOR + 4


def open_disk(path):
    with open(path, 'rb') as f:
        data = f.read()
    ext = os.path.splitext(path)[1].lower()
    if not data:  # an empty file: a new disk, by its extension
        if ext == '.scl':
            return SCL.blank()
        if ext == '.trd':
            return TRD.blank()
        raise Error('an empty file: name it .trd or .scl to start a disk')
    trd = len(data) >= 9 * SECTOR and data[TRDOS_ID] == 0x10
    if data[:8] == b'SINCLAIR' and len(data) > 8 and (ext == '.scl' or scl_size_ok(data) or not trd):
        return SCL(data)  # a TRD's first file may be named SINCLAIR too: the extension or an SCL's size tells
    return TRD(data)


def save(path, data):
    """Writes the image anew through a temporary file, so a failure leaves the old one whole; in place when the
    image has other hard links (a new file would part them) or the folder takes no new files."""
    path = os.path.realpath(path)
    st = os.stat(path)
    if not os.access(path, os.W_OK):
        raise Error('%s: the image is read-only' % path)
    fd = None
    if st.st_nlink == 1:
        try:
            fd, tmp = tempfile.mkstemp(dir=os.path.dirname(path), prefix='.' + os.path.basename(path) + '.')
        except OSError:  # no right to the folder
            pass
    if fd is None:
        with open(path, 'r+b') as f:
            f.write(data)
            f.truncate()
        return
    try:
        with os.fdopen(fd, 'wb') as f:
            f.write(data)
        os.chmod(tmp, stat.S_IMODE(st.st_mode))
        os.replace(tmp, path)
    except BaseException:
        os.unlink(tmp)
        raise


# Names. A TR-DOS name is 8 bytes padded with spaces; mc's name is the name, a dot and the type: "boot.B".
# The bytes that can't stand in a file name as they are (/, %, the controls, #7F, #FF, a leading space) are %XX;
# #80..#FE are cp866 (Russian disks name files so). The reverse gives the same bytes back.

def show_char(c, first=False):
    if c < 0x20 or c in (0x25, 0x2F, 0x7F, 0xFF) or c == 0x20 and first:
        return '%%%02X' % c
    if c >= 0x80:
        return bytes([c]).decode('cp866')
    return chr(c)


def show_text(b):
    return ''.join(show_char(c) for c in b)


def show_name(name):
    s = name.rstrip(b' ')
    return ''.join(show_char(c, i == 0) for i, c in enumerate(s)) or '%20'


def show_type(c):
    """A byte of the extension: a dot is %2E there too (the extension is what follows the last dot)."""
    return '%2E' if c == 0x2E else show_char(c, True)


def show_ext(e):
    lo, hi = e.p1 & 0xFF, e.p1 >> 8
    if e.type in EXT3 and lo in EXT3 and hi in EXT3:
        return chr(e.type) + chr(lo) + chr(hi)
    if e.type in EXT3 and lo == 0x5F and 0x21 <= hi <= 0x7E:  # QC's packed file: "_" and the original type
        return chr(e.type) + '_' + show_type(hi)
    return show_type(e.type)


def read_char(ch):
    if ord(ch) < 0x80:
        return ord(ch)
    try:
        return ch.encode('cp866')[0]
    except UnicodeEncodeError:
        return ord('_')


def read_text(s):
    out = bytearray()
    i = 0
    while i < len(s):
        if s[i] == '%' and re.match(r'[0-9A-Fa-f]{2}$', s[i + 1:i + 3]):
            out.append(int(s[i + 1:i + 3], 16))
            i += 3
        else:
            out.append(read_char(s[i]))
            i += 1
    return bytes(out)


def raw_name(e):
    return show_name(e.name) + '.' + show_ext(e)


def hobeta_name(e):
    return show_name(e.name) + '.$' + show_type(e.type)


def unique(files, name):
    """The repeated names (TR-DOS allows them) get ~N, N the file's number in the catalog (CATALOG's N): "x.C",
    "x.C~5". A TRD's rm only marks the file, so the numbers stay as mc saw them."""
    seen, out = set(), []
    for e in files:
        n = name(e)
        out.append(n if n not in seen else '%s~%d' % (n, e.index + 1))
        seen.add(n)
    return out


def listing(disk):
    """The root's and hobeta/'s names of the disk's files: two dicts of name -> Entry, in the catalog's order."""
    files = disk.files()
    raw = unique(files, raw_name)
    hob = unique(files, hobeta_name)
    return dict(zip(raw, files)), dict(zip(hob, files)), raw, hob


def parse_name(fname):
    """mc's name of a file to put on the disk -> (name, type, start or None, autostart line or None).
    "name.C", "name.$C" (hobeta), "name.pt3" (type p, start "t3"), "name.C@24576" (or @#6000, @0x6000; for
    BASIC, the autostart line); a repeat's ~2 is dropped; no extension: type C."""
    at = None
    m = re.match(r'(.+)@(0[xX][0-9A-Fa-f]+|#[0-9A-Fa-f]+|\d+)$', fname)
    if m:
        fname, num = m.group(1), m.group(2)
        at = int(num[1:], 16) if num[0] == '#' else int(num, 0) if num[:2].lower() == '0x' else int(num)
        if at > 0xFFFF:
            raise Error('%s: the number is past 65535' % num)
    stem, dot, ext = fname.rpartition('.')
    if not dot or not stem:
        stem, ext = fname, ''
    ext = re.sub(r'~\d+$', '', ext)
    if ext.startswith('$'):
        ext = ext[1:]
    name = (read_text(stem)[:8] + b' ' * 8)[:8]
    ext = read_text(ext)
    start = None
    if len(ext) == 3:
        start = ext[1] | ext[2] << 8
    return name, ext[0] if ext else ord('C'), start, at


# Contents.

def raw_length(e):
    """The file's length on the PC: the catalog's when it fits the sectors, else all the sectors."""
    full = e.sectors * SECTOR
    if e.type == ord('B'):  # the program with the variables, then #80, #AA and the autostart line
        n = e.p1
        if (n + SECTOR - 1) // SECTOR == e.sectors or (n + 4 + SECTOR - 1) // SECTOR == e.sectors:
            return n
        return full
    if (e.p2 + SECTOR - 1) // SECTOR == e.sectors:
        return e.p2
    return full


def hobeta(e, body):
    h = e.header() + bytes([0, e.sectors])
    return h + struct.pack('<H', (105 + 257 * sum(h)) & 0xFFFF) + body


def parse_hobeta(data):
    """A hobeta file -> (Entry, body in whole sectors), or None if the header's checksum is wrong."""
    if len(data) < HOBETA_HEADER:
        return None
    h = data[:15]
    if u16(data, 15) != (105 + 257 * sum(h)) & 0xFFFF:
        return None
    # Byte 14 is the sectors (13-14: their length in bytes); some old writers put the sectors in byte 13, so #00NN
    # is NN sectors when there is more than one sector of data, else NN bytes.
    if h[13] == 0:
        sectors = h[14]
    elif h[14] == 0:
        sectors = h[13] if len(data) - HOBETA_HEADER > SECTOR else 1
    else:
        sectors = (h[13] + 256 * h[14] + SECTOR - 1) // SECTOR
    if sectors == 0 and len(data) > HOBETA_HEADER:
        sectors = (len(data) - HOBETA_HEADER + SECTOR - 1) // SECTOR
    if sectors > MAX_SECTORS:
        return None
    body = bytes(data[HOBETA_HEADER:HOBETA_HEADER + sectors * SECTOR])
    p1, p2 = struct.unpack('<HH', bytes(h[9:13]))
    return Entry(bytes(h[0:8]), h[8], p1, p2, sectors), body + bytes(sectors * SECTOR - len(body))


def sectors_of(data):
    n = (len(data) + SECTOR - 1) // SECTOR
    if n > MAX_SECTORS:
        raise Error('%d bytes: a TR-DOS file holds %d at most' % (len(data), MAX_SECTORS * SECTOR))
    return data + bytes(n * SECTOR - len(data))


def new_file(disk, fname, data, old=None):
    """The Entry and the body to put on the disk for mc's name fname and the file's bytes; old is the file of that
    name it overwrites (mc's F4 saves so), or None."""
    e, body = new_entry(disk, fname, data, old)
    if e.name[0] in (0, 1):  # #00 ends the catalog, #01 is a deleted file
        raise Error('%s: a name starting with #%02X would not be seen on the disk' % (fname, e.name[0]))
    return e, body


def new_entry(disk, fname, data, old):
    name, type_, start, at = parse_name(fname)
    h = parse_hobeta(data)
    if h:
        e, body = h
        # A copy within the disk (F6 renaming, F5 under another name): the name is mc's.
        if any(x.same(e) and disk.body(x) == body for x in disk.files()):
            e.name, e.type = name, type_
        return e, body
    # A raw file: the same bytes on the disk already (a copy within it) give their start, length and sectors.
    for x in disk.files():
        b = disk.body(x)
        if b[:raw_length(x)] == data:
            if at is None or type_ != ord('B'):
                p1 = at if at is not None else start if start is not None else x.p1
                return Entry(name, type_, p1, x.p2, x.sectors), b
    if type_ == ord('B'):  # a BASIC program: no variables known; @N is the autostart line
        line, p2 = at, len(data)
        if old is not None:  # an edited program keeps the old one's autostart, and its variables if as long
            if at is None:
                line = autostart(old, disk.body(old))
            if len(data) == old.p1 and old.p2 <= len(data):
                p2 = old.p2
        tail = b'\x80\xaa' + struct.pack('<H', line) if line is not None else b''
        return Entry(name, type_, len(data), p2, 0), sectors_of(data + tail)
    if at is not None:
        start = at
    elif start is None and old is not None:
        start = old.p1  # an edited file keeps its start
    elif start is None:
        start = 16384 if len(data) == 6912 else 32768  # a screen, or code
    return Entry(name, type_, start, len(data), 0), sectors_of(data)


def autostart(e, body):
    """A BASIC program's autostart line: #80, #AA and the line after the program with the variables; or None."""
    if e.type == ord('B') and len(body) >= e.p1 + 4 and body[e.p1:e.p1 + 2] == b'\x80\xaa':
        return u16(body, e.p1 + 2)
    return None


def catalog_text(disk):
    lines = disk.describe() + ['', '  N  Name      Type  Start          Length  Sectors  Track:Sector']
    for i, e in enumerate(disk.entries):
        note = []
        line = autostart(e, disk.body(e))
        if line is not None:
            note.append('LINE %d' % line)
        if e.deleted:
            note.append('deleted')
        place = '%d:%d' % (e.start // 16, e.start % 16) if disk.kind == 'TRD' else ''
        name = (show_char(e.name[0] if not e.deleted else 0x3F, True) + show_text(e.name[1:])).ljust(8)
        lines.append('%3d  %s  %-4s  %-5d  #%04X  %-6d  %-7d  %-12s  %s' % (
            i + 1, name, show_char(e.type, True), e.p1, e.p1, e.p2, e.sectors, place, ', '.join(note)))
    return ('\n'.join(line.rstrip() for line in lines) + '\n').encode('utf-8')


# The commands.

def split(stored):
    """mc's path in the archive -> ('' or 'hobeta', the file's name)."""
    parts = [p for p in stored.split('/') if p not in ('', '.')]
    if len(parts) == 2 and parts[0] == HOBETA_DIR:
        return HOBETA_DIR, parts[1]
    if len(parts) == 1:
        return '', parts[0]
    raise Error('%s: TR-DOS has no folders' % stored)


def lookup(disk, folder, name):
    """The file of mc's name, or None. A repeat's "x.C~N" is the file N of the catalog while it is an x.C: mc doesn't
    reread the listing after rm, and the name of the next x.C may have changed meanwhile."""
    raw, hob, _, _ = listing(disk)
    e = (hob if folder else raw).get(name)
    m = re.match(r'(.+)~(\d+)$', name)
    if e is None and m and 0 < int(m.group(2)) <= len(disk.entries):
        x = disk.entries[int(m.group(2)) - 1]
        if not x.deleted and (hobeta_name if folder else raw_name)(x) == m.group(1):
            e = x
    return e


def find(disk, stored):
    folder, name = split(stored)
    e = lookup(disk, folder, name)
    if e is None:
        raise Error('%s: no such file on the disk' % stored)
    return folder, e


def cmd_list(path):
    disk = open_disk(path)
    date = time.strftime('%m-%d-%Y %H:%M', time.localtime(os.stat(path).st_mtime))
    owner = '%d %d' % (os.getuid(), os.getgid())
    raw, hob, raw_names, hob_names = listing(disk)
    lines = ['-r--r--r-- 1 %s %d %s %s' % (owner, len(catalog_text(disk)), date, CATALOG),
             'drwxr-xr-x 1 %s 0 %s %s' % (owner, date, HOBETA_DIR)]
    for n in raw_names:
        lines.append('-rw-r--r-- 1 %s %d %s %s' % (owner, raw_length(raw[n]), date, n))
    for n in hob_names:
        lines.append('-rw-r--r-- 1 %s %d %s %s/%s' % (owner, HOBETA_HEADER + hob[n].sectors * SECTOR, date,
                                                        HOBETA_DIR, n))
    return ''.join(line + '\n' for line in lines)


def cmd_copyout(path, stored, dest):
    disk = open_disk(path)
    if split(stored) == ('', CATALOG):
        data = catalog_text(disk)
    else:
        folder, e = find(disk, stored)
        body = disk.body(e)
        data = hobeta(e, body) if folder else body[:raw_length(e)]
    with open(dest, 'wb') as f:
        f.write(data)


def cmd_copyin(path, stored, src):
    disk = open_disk(path)
    folder, name = split(stored)
    if not folder and name in (CATALOG, HOBETA_DIR):
        raise Error('%s is not a file of the disk' % name)
    with open(src, 'rb') as f:
        data = f.read()
    old = lookup(disk, folder, name)
    new, body = new_file(disk, name, data, old)
    if old is not None:
        disk.replace(old, new, body)
    else:
        disk.add(new, body)
    save(path, disk.bytes())


def cmd_rm(path, stored):
    disk = open_disk(path)
    if split(stored) == ('', CATALOG):
        raise Error('%s is not a file of the disk' % CATALOG)
    _, e = find(disk, stored)
    disk.delete(e)
    save(path, disk.bytes())


def main(argv):
    # mc gives the names back as the list wrote them (UTF-8): take argv's bytes as they came.
    args = [os.fsencode(a).decode('utf-8', 'surrogateescape') for a in argv[1:]]
    cmd = args[0] if args else ''
    n = {'list': 2, 'catalog': 2, 'copyout': 4, 'copyin': 4, 'rm': 3, 'mkdir': 3, 'rmdir': 3}
    if cmd not in n or len(args) != n[cmd]:
        sys.stderr.write('usage: trdos list|catalog ARCHIVE | copyout|copyin ARCHIVE NAME FILE | rm ARCHIVE NAME\n')
        return 2
    path = argv[2]
    try:
        if cmd == 'list':
            sys.stdout.buffer.write(cmd_list(path).encode('utf-8', 'surrogateescape'))
        elif cmd == 'catalog':
            sys.stdout.buffer.write(catalog_text(open_disk(path)))
        elif cmd == 'copyout':
            cmd_copyout(path, args[2], argv[4])
        elif cmd == 'copyin':
            cmd_copyin(path, args[2], argv[4])
        elif cmd == 'rm':
            cmd_rm(path, args[2])
        else:
            raise Error('TR-DOS has no folders')
    except (Error, OSError) as e:
        sys.stderr.write('trdos: %s\n' % e)
        return 1
    return 0


if __name__ == '__main__':
    sys.exit(main(sys.argv))
