#!/usr/bin/python3
"""Polkit-authorized ECZOS network optical drive operations."""

import argparse
import fcntl
import hashlib
import html
import ipaddress
import json
import logging
import logging.handlers
import os
import re
import socket
import subprocess
import sys
import tempfile
import time
from datetime import datetime, timezone
from pathlib import Path

STATE_DIR = Path("/var/lib/eczos/network-optical")
STATE = STATE_DIR / "state.json"
LOCK = STATE_DIR / "state.lock"
AVAHI_DIR = Path("/etc/avahi/services")
UDEV_RULES_DIR = Path("/etc/udev/rules.d")
CONFIG = Path("/etc/eczos/network-optical.conf")
IQN_PATTERN = re.compile(r"iqn\.[A-Za-z0-9.:-]{8,223}")
HOST_PATTERN = re.compile(r"[A-Za-z0-9](?:[A-Za-z0-9.:-]{0,251}[A-Za-z0-9])?")


def logger():
    log = logging.getLogger("eczos-network-optical")
    log.setLevel(logging.INFO)
    if not log.handlers:
        try:
            log.addHandler(logging.handlers.SysLogHandler(address="/dev/log"))
        except OSError:
            log.addHandler(logging.StreamHandler(sys.stderr))
    return log


LOG = logger()


def fail(message, code=1):
    print(message, file=sys.stderr)
    raise SystemExit(code)


def run(argv, timeout=30, check=True):
    try:
        result = subprocess.run(argv, text=True, capture_output=True, timeout=timeout, check=False)
    except (OSError, subprocess.TimeoutExpired) as error:
        fail(f"Backend command unavailable: {error}")
    if check and result.returncode != 0:
        detail = (result.stderr or result.stdout).strip().splitlines()
        fail(detail[-1] if detail else "The backend operation failed.")
    return result


def set_optical_firewall_service(enabled):
    """Keep TCP 3260 open only while this computer exports a drive."""
    service = "eczos-network-optical"
    zone = "eczos-public"
    firewall_cmd = Path("/usr/bin/firewall-cmd")
    offline_cmd = Path("/usr/bin/firewall-offline-cmd")
    if not firewall_cmd.exists() and not offline_cmd.exists():
        return True

    action = "--add-service" if enabled else "--remove-service"
    commands = []
    running = False
    if firewall_cmd.exists():
        state = subprocess.run(
            [str(firewall_cmd), "--state"], text=True, capture_output=True,
            timeout=10, check=False)
        running = state.returncode == 0
    if running:
        commands.extend([
            [str(firewall_cmd), f"--zone={zone}", f"{action}={service}"],
            [str(firewall_cmd), "--permanent", f"--zone={zone}", f"{action}={service}"],
        ])
    elif offline_cmd.exists():
        commands.append([str(offline_cmd), f"--zone={zone}", f"{action}={service}"])
    else:
        return False

    for command in commands:
        result = subprocess.run(command, text=True, capture_output=True, timeout=15, check=False)
        detail = f"{result.stdout}\n{result.stderr}".lower()
        if result.returncode != 0 and not any(
                marker in detail for marker in ("already_enabled", "already disabled", "not_enabled")):
            LOG.error("firewall update failed: %s", (result.stderr or result.stdout).strip())
            return False
    return True


def config():
    values = {"PortalAddress": "", "Port": "3260", "ReleaseTimeoutSec": "30"}
    try:
        for line in CONFIG.read_text(encoding="utf-8").splitlines():
            key, separator, value = line.partition("=")
            if separator and key in values:
                values[key] = value.strip()
    except OSError:
        pass
    return values


def locked_state():
    STATE_DIR.mkdir(mode=0o755, parents=True, exist_ok=True)
    os.chmod(STATE_DIR, 0o755)
    descriptor = LOCK.open("a+")
    os.chmod(LOCK, 0o600)
    fcntl.flock(descriptor, fcntl.LOCK_EX)
    try:
        data = json.loads(STATE.read_text(encoding="utf-8"))
        if data.get("schema") != 1 or not isinstance(data.get("shares"), dict):
            raise ValueError
    except (OSError, ValueError):
        data = {"schema": 1, "shares": {}, "clients": {}}
    data.setdefault("clients", {})
    return descriptor, data


def save_state(descriptor, data):
    fd, temporary = tempfile.mkstemp(prefix="state.", dir=STATE_DIR)
    try:
        with os.fdopen(fd, "w", encoding="utf-8") as stream:
            json.dump(data, stream, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
            stream.write("\n")
            stream.flush()
            os.fsync(stream.fileno())
        os.chmod(temporary, 0o644)
        os.replace(temporary, STATE)
    finally:
        if os.path.exists(temporary):
            os.unlink(temporary)
        fcntl.flock(descriptor, fcntl.LOCK_UN)
        descriptor.close()


def release_state(descriptor):
    fcntl.flock(descriptor, fcntl.LOCK_UN)
    descriptor.close()


def properties(device):
    result = run(["/usr/bin/udevadm", "info", "--query=property", "--name", device])
    values = {}
    for line in result.stdout.splitlines():
        key, separator, value = line.partition("=")
        if separator:
            values[key] = value
    return values


def local_drive(device):
    resolved = os.path.realpath(device)
    if not re.fullmatch(r"/dev/sr[0-9]+", resolved) or not Path(resolved).is_block_device():
        fail("Select a physical optical drive.")
    props = properties(resolved)
    transport = run(["/usr/bin/lsblk", "-dnro", "TRAN", "--", resolved]).stdout.strip()
    if props.get("ID_CDROM") != "1" or transport == "iscsi" or props.get("ID_BUS") == "iscsi":
        fail("Only a local physical optical drive can be shared.")
    stable = props.get("ID_SERIAL") or props.get("ID_PATH") or str(Path(f"/sys/class/block/{Path(resolved).name}").resolve())
    device_id = hashlib.sha256(stable.encode()).hexdigest()[:16]
    vendor = (props.get("ID_VENDOR") or "").replace("_", " ").strip()
    model = (props.get("ID_MODEL") or "Optical drive").replace("_", " ").strip()
    write = any(props.get(key) == "1" for key in (
        "ID_CDROM_CD_R", "ID_CDROM_CD_RW", "ID_CDROM_DVD_R", "ID_CDROM_DVD_RW",
        "ID_CDROM_DVD_PLUS_R", "ID_CDROM_DVD_PLUS_RW", "ID_CDROM_DVD_RAM",
        "ID_CDROM_BD_R", "ID_CDROM_BD_RE"))
    return resolved, props, device_id, vendor, model, write


def lan_address(requested=""):
    candidates = []
    if requested:
        candidates.append(requested)
    configured = config().get("PortalAddress", "")
    if configured:
        candidates.append(configured)
    result = run(["/usr/sbin/ip", "-j", "address", "show", "up"])
    for interface in json.loads(result.stdout):
        if interface.get("ifname") == "lo":
            continue
        for info in interface.get("addr_info", []):
            if info.get("scope") == "global":
                candidates.append(info.get("local", ""))
    for value in candidates:
        try:
            address = ipaddress.ip_address(value)
        except ValueError:
            continue
        if address.is_private and not address.is_loopback and not address.is_unspecified:
            return str(address)
    fail("No active private LAN address is available.")


def machine_token():
    value = Path("/etc/machine-id").read_text().strip().lower()
    if not re.fullmatch(r"[0-9a-f]{32}", value):
        fail("The stable machine identity is unavailable.")
    return value[:16]


def targetctl_save():
    executable = "/usr/bin/targetctl"
    if not Path(executable).exists():
        executable = "/usr/sbin/targetctl"
    run([executable, "save"])


def avahi_path(device_id):
    return AVAHI_DIR / f"eczos-optical-{device_id}.service"


def udev_rule_path(device_id):
    return UDEV_RULES_DIR / f"99-eczos-network-optical-{device_id}.rules"


def set_local_use_blocked(device_id, device, blocked, match_key="", match_value=""):
    path = udev_rule_path(device_id)
    if blocked:
        UDEV_RULES_DIR.mkdir(parents=True, exist_ok=True)
        if match_key not in ("ID_SERIAL", "ID_PATH") or not match_value:
            match_key, match_value = "DEVNAME", device
        if any(character in match_value for character in ('\n', '\r', '\x00')):
            fail("The optical drive identity is invalid.")
        escaped = match_value.replace("\\", "\\\\").replace('"', '\\"')
        path.write_text(
            f'# Managed by ECZOS Network Optical Drives\nENV{{{match_key}}}=="{escaped}", ENV{{UDISKS_IGNORE}}="1"\n',
            encoding="utf-8")
        os.chmod(path, 0o644)
    elif path.exists():
        path.unlink()
    run(["/usr/bin/udevadm", "control", "--reload"], check=False)
    run(["/usr/bin/udevadm", "trigger", "--action=change", "--name-match", Path(device).name], check=False)


def write_avahi(share):
    AVAHI_DIR.mkdir(parents=True, exist_ok=True)
    txt = {
        "version": "1", "hostname": socket.gethostname(), "vendor": share["vendor"],
        "model": share["model"], "iqn": share["iqn"],
        "write": "true" if share["write"] else "false",
        "cd": "true", "dvd": "true", "bluray": "true" if share.get("bluray") else "false",
        "capabilities": share["capabilities"], "exclusive": "dynamic-acl-v1",
    }
    lines = ["<?xml version=\"1.0\" standalone='no'?>", "<!DOCTYPE service-group SYSTEM \"avahi-service.dtd\">",
             "<service-group>", f"  <name replace-wildcards=\"yes\">{html.escape(share['name'])}</name>",
             "  <service>", "    <type>_eczos-optical._tcp</type>", f"    <port>{share['port']}</port>"]
    lines.extend(f"    <txt-record>{html.escape(key + '=' + str(value))}</txt-record>" for key, value in txt.items())
    lines.extend(["  </service>", "</service-group>", ""])
    path = avahi_path(share["id"])
    path.write_text("\n".join(lines), encoding="utf-8")
    os.chmod(path, 0o644)


def capability_summary(props):
    values = []
    if any(props.get(key) == "1" for key in ("ID_CDROM_BD", "ID_CDROM_BD_R", "ID_CDROM_BD_RE")):
        values.append("Blu-ray")
    if any(props.get(key) == "1" for key in ("ID_CDROM_DVD_R", "ID_CDROM_DVD_RW", "ID_CDROM_DVD_PLUS_R", "ID_CDROM_DVD_PLUS_RW", "ID_CDROM_DVD_RAM")):
        values.append("DVD±RW")
    elif props.get("ID_CDROM_DVD") == "1":
        values.append("DVD-ROM")
    values.append("CD-RW" if any(props.get(key) == "1" for key in ("ID_CDROM_CD_R", "ID_CDROM_CD_RW")) else "CD-ROM")
    return " / ".join(values)


def share(args):
    device, props, device_id, vendor, model, writable = local_drive(args.device)
    if len(args.name) > 80 or any(ord(char) < 32 for char in args.name):
        fail("The share name is invalid.")
    address = lan_address(args.address)
    mounts = run(["/usr/bin/findmnt", "-rn", "-S", device], check=False)
    if mounts and mounts.stdout.strip():
        run(["/usr/bin/udisksctl", "unmount", "--block-device", device])
    busy = run(["/usr/bin/fuser", "--", device], check=False)
    if busy and busy.returncode == 0:
        fail("The drive is currently in use. Stop reading or writing before sharing it.", 4)
    lock, state = locked_state()
    if device_id in state["shares"]:
        release_state(lock)
        fail("This drive is already shared.")
    token = machine_token()
    month = datetime.now(timezone.utc).strftime("%Y-%m")
    iqn = f"iqn.{month}.nl.easycomp.eczos:{token}.optical.{device_id[:12]}"
    backstore = f"eczos_{device_id[:12]}"
    storage = None
    target = None
    try:
        from rtslib_fb import FabricModule, LUN, NetworkPortal, PSCSIStorageObject, TPG, Target
        storage = PSCSIStorageObject(backstore, dev=device)
        target = Target(FabricModule("iscsi"), iqn, mode="create")
        tpg = TPG(target, 1, mode="create")
        LUN(tpg, 0, storage_object=storage)
        tpg.set_attribute("authentication", "0")
        tpg.set_attribute("generate_node_acls", "1")
        tpg.set_attribute("cache_dynamic_acls", "1")
        tpg.set_attribute("demo_mode_write_protect", "0")
        NetworkPortal(tpg, address, 3260)
        tpg.enable = True
        targetctl_save()
        if not set_optical_firewall_service(True):
            raise RuntimeError("the ECZOS firewall rule could not be enabled")
    except Exception as error:
        try:
            if target is not None:
                target.delete()
            if storage is not None:
                storage.delete()
            targetctl_save()
        except Exception as cleanup_error:
            LOG.error("share rollback failed for %s: %s", device_id, cleanup_error)
        if not state["shares"]:
            set_optical_firewall_service(False)
        release_state(lock)
        LOG.error("share failed for %s: %s", device_id, error)
        fail(f"The optical drive could not be shared: {error}")
    record = {
        "id": device_id, "device": device, "name": args.name, "vendor": vendor,
        "model": model, "capabilities": capability_summary(props), "write": writable,
        "bluray": any(props.get(key) == "1" for key in ("ID_CDROM_BD", "ID_CDROM_BD_R", "ID_CDROM_BD_RE")),
        "iqn": iqn, "backstore": backstore, "address": address, "port": 3260,
        "claim": "", "claimMissingSince": 0,
        "udevMatchKey": "ID_SERIAL" if props.get("ID_SERIAL") else ("ID_PATH" if props.get("ID_PATH") else "DEVNAME"),
        "udevMatchValue": props.get("ID_SERIAL") or props.get("ID_PATH") or device,
    }
    state["shares"][device_id] = record
    save_state(lock, state)
    set_local_use_blocked(device_id, device, True, record["udevMatchKey"], record["udevMatchValue"])
    write_avahi(record)
    run(["/usr/bin/systemctl", "enable", "--now", "eczos-network-optical-guard.service"])
    LOG.info("shared %s as %s on %s", device, iqn, address)
    print(json.dumps({"ok": True, "share": record}, separators=(",", ":")))


def unshare(args):
    lock, state = locked_state()
    record = state["shares"].get(args.id)
    if not record:
        release_state(lock)
        fail("The selected share no longer exists.")
    try:
        from rtslib_fb import FabricModule, PSCSIStorageObject, Target
        target = Target(FabricModule("iscsi"), record["iqn"], mode="lookup")
        if any(acl.session for tpg in target.tpgs for acl in tpg.node_acls) and not args.force:
            release_state(lock)
            fail("The drive is in use. Disconnect the client first.", 4)
        target.delete()
        try:
            PSCSIStorageObject(record["backstore"]).delete()
        except Exception:
            pass
        targetctl_save()
    except SystemExit:
        raise
    except Exception as error:
        release_state(lock)
        LOG.error("unshare failed for %s: %s", args.id, error)
        fail(f"Sharing could not be stopped: {error}")
    path = avahi_path(args.id)
    if path.exists():
        path.unlink()
    set_local_use_blocked(args.id, record["device"], False)
    del state["shares"][args.id]
    save_state(lock, state)
    if not state["shares"] and not set_optical_firewall_service(False):
        LOG.warning("the optical-drive firewall rule could not be removed")
    LOG.info("stopped share %s", record["iqn"])
    print(json.dumps({"ok": True}))


def endpoint(host, port):
    if port != 3260:
        fail("The discovered server address is invalid.")
    try:
        address = ipaddress.ip_address(host)
        if not address.is_private:
            fail("Only private LAN servers are allowed.")
    except ValueError:
        if not HOST_PATTERN.fullmatch(host) or host.startswith("-") or not host.endswith(".local"):
            fail("Only local-network hostnames are allowed.")
    return f"[{host}]:{port}" if ":" in host else f"{host}:{port}"


def validate_target(iqn):
    if not IQN_PATTERN.fullmatch(iqn):
        fail("The discovered target identifier is invalid.")


def connect(args):
    validate_target(args.iqn)
    portal = endpoint(args.host, args.port)
    discovery = run(["/usr/sbin/iscsiadm", "-m", "discovery", "-t", "sendtargets", "-p", portal])
    if not any(args.iqn in line for line in discovery.stdout.splitlines()):
        fail("The selected optical drive was not returned by the server.")
    startup = "automatic" if args.auto else "manual"
    run(["/usr/sbin/iscsiadm", "-m", "node", "-T", args.iqn, "-p", portal,
         "--op", "update", "-n", "node.startup", "-v", startup])
    result = run(["/usr/sbin/iscsiadm", "-m", "node", "-T", args.iqn, "-p", portal, "--login"], check=False)
    if result.returncode != 0:
        detail = (result.stderr or result.stdout).lower()
        if "authorization" in detail or "authentication" in detail or "access" in detail:
            fail("The drive is in use by another computer or authentication was refused.", 5)
        fail("Connecting to the network optical drive failed.")
    run(["/usr/bin/udevadm", "settle", "--timeout=15"], timeout=20, check=False)
    lock, state = locked_state()
    state["clients"][args.iqn] = {
        "host": args.host, "port": args.port, "automatic": bool(args.auto)
    }
    save_state(lock, state)
    LOG.info("connected target %s at %s", args.iqn, portal)
    print(json.dumps({"ok": True, "iqn": args.iqn}))


def target_devices(iqn):
    devices = []
    for session in Path("/sys/class/iscsi_session").glob("session*"):
        try:
            if (session / "targetname").read_text().strip() != iqn:
                continue
        except OSError:
            continue
        devices.extend(f"/dev/{path.name}" for path in session.glob("device/target*/*/block/sr*"))
    return devices


def disconnect(args):
    validate_target(args.iqn)
    portal = endpoint(args.host, args.port)
    devices = target_devices(args.iqn)
    for device in devices:
        busy = run(["/usr/bin/fuser", "--", device], check=False)
        if busy and busy.returncode == 0 and not args.force:
            fail("The drive is currently in use. Stop reading or writing first.", 4)
        mounts = run(["/usr/bin/findmnt", "-rn", "-S", device], check=False)
        if mounts and mounts.stdout.strip():
            run(["/usr/bin/udisksctl", "unmount", "--block-device", device])
    run(["/usr/bin/sync"])
    run(["/usr/sbin/iscsiadm", "-m", "node", "-T", args.iqn, "-p", portal, "--logout"])
    for _ in range(50):
        if not target_devices(args.iqn):
            break
        time.sleep(0.1)
    if target_devices(args.iqn):
        fail("The remote optical device did not disappear cleanly.")
    LOG.info("disconnected target %s", args.iqn)
    print(json.dumps({"ok": True}))


def set_auto(args):
    validate_target(args.iqn)
    portal = endpoint(args.host, args.port)
    run(["/usr/sbin/iscsiadm", "-m", "node", "-T", args.iqn, "-p", portal,
         "--op", "update", "-n", "node.startup", "-v", "automatic" if args.enabled else "manual"])
    lock, state = locked_state()
    state["clients"][args.iqn] = {
        "host": args.host, "port": args.port, "automatic": bool(args.enabled)
    }
    save_state(lock, state)
    print(json.dumps({"ok": True}))


def forget(args):
    validate_target(args.iqn)
    portal = endpoint(args.host, args.port)
    if target_devices(args.iqn):
        fail("Disconnect the drive before forgetting it.")
    run(["/usr/sbin/iscsiadm", "-m", "node", "-T", args.iqn, "-p", portal, "--op", "delete"])
    lock, state = locked_state()
    state["clients"].pop(args.iqn, None)
    save_state(lock, state)
    print(json.dumps({"ok": True}))


def eject_media(args):
    device = os.path.realpath(args.device)
    if not re.fullmatch(r"/dev/sr[0-9]+", device) or not Path(device).is_block_device():
        fail("The selected optical device is invalid.")
    busy = run(["/usr/bin/fuser", "--", device], check=False)
    if busy and busy.returncode == 0:
        fail("The drive is currently in use.", 4)
    mounts = run(["/usr/bin/findmnt", "-rn", "-S", device], check=False)
    if mounts and mounts.stdout.strip():
        run(["/usr/bin/udisksctl", "unmount", "--block-device", device])
    run(["/usr/bin/eject", device])
    print(json.dumps({"ok": True}))


def parser():
    root = argparse.ArgumentParser(prog="eczos-network-optical-helper")
    commands = root.add_subparsers(dest="command", required=True)
    share_parser = commands.add_parser("share")
    share_parser.add_argument("--device", required=True)
    share_parser.add_argument("--name", required=True)
    share_parser.add_argument("--address", default="")
    unshare_parser = commands.add_parser("unshare")
    unshare_parser.add_argument("--id", required=True)
    unshare_parser.add_argument("--force", action="store_true")
    for command in ("connect", "disconnect", "set-auto", "forget"):
        item = commands.add_parser(command)
        item.add_argument("--host", required=True)
        item.add_argument("--port", type=int, default=3260)
        item.add_argument("--iqn", required=True)
        if command == "connect":
            item.add_argument("--auto", action="store_true")
        if command == "disconnect":
            item.add_argument("--force", action="store_true")
        if command == "set-auto":
            item.add_argument("--enabled", action="store_true")
    eject_parser = commands.add_parser("eject")
    eject_parser.add_argument("--device", required=True)
    return root


def main():
    if os.geteuid() != 0:
        fail("This helper must be started through system authentication.", 2)
    args = parser().parse_args()
    {"share": share, "unshare": unshare, "connect": connect, "disconnect": disconnect,
     "set-auto": set_auto, "forget": forget, "eject": eject_media}[args.command](args)


if __name__ == "__main__":
    main()
