#!/usr/bin/python3
"""Read-only ECZOS network optical drive discovery and status CLI."""

import argparse
import hashlib
import ipaddress
import json
import os
import re
import shlex
import socket
import subprocess
from pathlib import Path

STATE = Path("/var/lib/eczos/network-optical/state.json")
SERVICE_TYPE = "_eczos-optical._tcp"


def run(argv, timeout=8):
    try:
        return subprocess.run(argv, text=True, capture_output=True, timeout=timeout, check=False)
    except (OSError, subprocess.TimeoutExpired):
        return None


def udev_properties(device):
    result = run(["/usr/bin/udevadm", "info", "--query=property", "--name", device])
    values = {}
    if result and result.returncode == 0:
        for line in result.stdout.splitlines():
            key, separator, value = line.partition("=")
            if separator and re.fullmatch(r"[A-Z0-9_]+", key):
                values[key] = value
    return values


def load_state():
    try:
        data = json.loads(STATE.read_text(encoding="utf-8"))
        if data.get("schema") == 1:
            data.setdefault("shares", {})
            data.setdefault("clients", {})
            return data
        return {"schema": 1, "shares": {}, "clients": {}}
    except (OSError, ValueError):
        return {"schema": 1, "shares": {}, "clients": {}}


def capability_summary(props):
    capabilities = []
    if props.get("ID_CDROM_BD") == "1" or any(props.get(key) == "1" for key in ("ID_CDROM_BD_R", "ID_CDROM_BD_RE")):
        capabilities.append("Blu-ray")
    dvd_write = 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_R_DL", "ID_CDROM_DVD_PLUS_R_DL",
        "ID_CDROM_DVD_RAM"))
    if dvd_write:
        capabilities.append("DVD±RW")
    elif props.get("ID_CDROM_DVD") == "1":
        capabilities.append("DVD-ROM")
    if props.get("ID_CDROM_CD_RW") == "1" or props.get("ID_CDROM_CD_R") == "1":
        capabilities.append("CD-RW")
    elif props.get("ID_CDROM") == "1":
        capabilities.append("CD-ROM")
    return " / ".join(capabilities) or "Optical drive"


def local_drives():
    state = load_state().get("shares", {})
    drives = []
    for block in sorted(Path("/sys/class/block").glob("sr*")):
        device = f"/dev/{block.name}"
        props = udev_properties(device)
        transport_result = run(["/usr/bin/lsblk", "-dnro", "TRAN", "--", device])
        transport = transport_result.stdout.strip() if transport_result and transport_result.returncode == 0 else ""
        if transport == "iscsi" or props.get("ID_BUS") == "iscsi":
            continue
        vendor = (props.get("ID_VENDOR") or props.get("ID_VENDOR_FROM_DATABASE") or "").replace("_", " ").strip()
        model = (props.get("ID_MODEL") or "Optical drive").replace("_", " ").strip()
        generic = ""
        generic_dir = block / "device/scsi_generic"
        if generic_dir.is_dir():
            entries = sorted(generic_dir.iterdir())
            if entries:
                generic = f"/dev/{entries[0].name}"
        stable_source = props.get("ID_SERIAL") or props.get("ID_PATH") or str(block.resolve())
        device_id = hashlib.sha256(stable_source.encode()).hexdigest()[:16]
        share = state.get(device_id, {})
        drives.append({
            "id": device_id,
            "device": device,
            "generic": generic,
            "vendor": vendor,
            "model": model,
            "label": " ".join(value for value in (vendor, model) if value).strip(),
            "capabilities": capability_summary(props),
            "media": props.get("ID_CDROM_MEDIA") == "1",
            "shared": bool(share),
            "shareName": share.get("name", ""),
            "iqn": share.get("iqn", ""),
            "address": share.get("address", ""),
            "port": share.get("port", 3260),
            "claim": share.get("claim", ""),
            "server": f"{socket.gethostname()}.local",
        })
    return drives


def unescape_avahi(value):
    decoded = re.sub(r"\\([0-9]{3})", lambda match: chr(int(match.group(1))), value)
    try:
        return decoded.encode("latin-1").decode("utf-8")
    except (UnicodeEncodeError, UnicodeDecodeError):
        return decoded


def sessions():
    found = {}
    for session in Path("/sys/class/iscsi_session").glob("session*"):
        try:
            iqn = (session / "targetname").read_text().strip()
        except OSError:
            continue
        devices = []
        for block in session.glob("device/target*/*/block/sr*"):
            devices.append(f"/dev/{block.name}")
        session_number = session.name.removeprefix("session")
        connections = sorted(Path("/sys/class/iscsi_connection").glob(f"connection{session_number}:*"))
        address = ""
        port = 3260
        if connections:
            try:
                address = (connections[0] / "persistent_address").read_text().strip()
                port = int((connections[0] / "persistent_port").read_text().strip())
            except (OSError, ValueError):
                try:
                    address = (connections[0] / "address").read_text().strip()
                    port = int((connections[0] / "port").read_text().strip())
                except (OSError, ValueError):
                    address, port = "", 3260
        found[iqn] = {"devices": devices, "address": address, "port": port}
    return found


def network_drives():
    result = run(["/usr/bin/avahi-browse", "--resolve", "--parsable", "--terminate", SERVICE_TYPE], timeout=12)
    active = sessions()
    automatic = {
        iqn for iqn, node in load_state().get("clients", {}).items()
        if isinstance(node, dict) and node.get("automatic") is True
    }
    records = {}
    if not result or result.returncode not in (0, 1):
        return []
    for line in result.stdout.splitlines():
        if not line.startswith("="):
            continue
        fields = line.split(";")
        if len(fields) < 10:
            continue
        name = unescape_avahi(fields[3])
        host = unescape_avahi(fields[6]).rstrip(".")
        address = fields[7]
        try:
            port = int(fields[8])
        except ValueError:
            continue
        txt = {}
        try:
            txt_records = shlex.split(";".join(fields[9:]))
        except ValueError:
            continue
        for item in txt_records:
            item = unescape_avahi(item)
            key, separator, value = item.partition("=")
            if separator and re.fullmatch(r"[a-z][a-z0-9_-]{0,31}", key):
                txt[key] = value[:512]
        iqn = txt.get("iqn", "")
        if txt.get("version") != "1" or port != 3260 or not re.fullmatch(r"iqn\.[A-Za-z0-9.:-]{8,223}", iqn):
            continue
        candidate = {
            "name": name,
            "host": host,
            "address": address,
            "port": port,
            "iqn": iqn,
            "vendor": txt.get("vendor", ""),
            "model": txt.get("model", "Optical drive"),
            "capabilities": txt.get("capabilities", "Optical drive"),
            "writable": txt.get("write") == "true",
            "connected": iqn in active,
            "automatic": iqn in automatic,
            "device": (active.get(iqn, {}).get("devices") or [""])[0],
            "status": "connected" if iqn in active else "available",
            "managed": True,
        }
        previous = records.get(iqn)
        try:
            candidate_is_v4 = ipaddress.ip_address(address).version == 4
            previous_is_v4 = bool(previous) and ipaddress.ip_address(previous["address"]).version == 4
        except ValueError:
            candidate_is_v4 = previous_is_v4 = False
        if not previous or candidate_is_v4 or not previous_is_v4:
            records[iqn] = candidate

    # A session created before ECZOS discovery support has no Avahi record. It
    # must still be visible while connected so the user is never left with an
    # invisible remote device that can only be controlled from a terminal.
    for iqn, session in active.items():
        if iqn in records or not session["devices"]:
            continue
        device = session["devices"][0]
        props = udev_properties(device)
        if props.get("ID_CDROM") != "1":
            continue
        vendor = (props.get("ID_VENDOR") or "").replace("_", " ").strip()
        model = (props.get("ID_MODEL") or "Optical drive").replace("_", " ").strip()
        records[iqn] = {
            "name": " ".join(value for value in (vendor, model) if value).strip(),
            "host": session["address"],
            "address": session["address"],
            "port": session["port"],
            "iqn": iqn,
            "vendor": vendor,
            "model": model,
            "capabilities": capability_summary(props),
            "writable": 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")),
            "connected": True,
            "automatic": iqn in automatic,
            "device": device,
            "status": "connected",
            "managed": False,
        }
    return sorted(records.values(), key=lambda item: (item["name"].casefold(), item["host"]))


def main():
    parser = argparse.ArgumentParser(prog="eczos-network-optical")
    parser.add_argument("command", choices=("list-local", "list-network", "status"))
    parser.add_argument("--json", action="store_true", required=True)
    args = parser.parse_args()
    payload = {"schema": 1}
    if args.command in ("list-local", "status"):
        payload["local"] = local_drives()
    if args.command in ("list-network", "status"):
        payload["network"] = network_drives()
    print(json.dumps(payload, ensure_ascii=False, separators=(",", ":")))


if __name__ == "__main__":
    main()
