#!/usr/bin/python3
"""Keep RustDesk-derived ECZHOA sessions away from crash-prone GPU encoders."""

from __future__ import annotations

import argparse
import json
import os
import pwd
import re
import stat
import subprocess
import sys
import tempfile
from datetime import datetime, timezone
from pathlib import Path


STATE_DIR = Path("/var/lib/eczos/remote-support")
STATE_FILE = STATE_DIR / "status.json"
SETTING = "enable-hwcodec"


def rustdesk_installed() -> bool:
    return Path("/usr/share/rustdesk/rustdesk").is_file() or Path("/usr/bin/rustdesk").exists()


def candidate_configs() -> list[tuple[Path, int, int]]:
    candidates: list[tuple[Path, int, int]] = []
    root = pwd.getpwnam("root")
    candidates.append((Path(root.pw_dir) / ".config/rustdesk/RustDesk2.toml", root.pw_uid, root.pw_gid))
    for entry in pwd.getpwall():
        if not (1000 <= entry.pw_uid < 60000) or not entry.pw_dir.startswith("/"):
            continue
        candidates.append((Path(entry.pw_dir) / ".config/rustdesk/RustDesk2.toml", entry.pw_uid, entry.pw_gid))
    try:
        sddm = pwd.getpwnam("sddm")
        candidates.append((Path(sddm.pw_dir) / ".config/rustdesk/RustDesk2.toml", sddm.pw_uid, sddm.pw_gid))
    except KeyError:
        pass
    return list(dict.fromkeys(candidates))


def software_codec_configured(path: Path) -> bool:
    try:
        if path.is_symlink() or not path.is_file():
            return False
        text = path.read_text(encoding="utf-8")
    except (OSError, UnicodeError):
        return False
    section = None
    for line in text.splitlines():
        match = re.match(r"^\s*\[([^]]+)]\s*$", line)
        if match:
            section = match.group(1)
        elif section == "options" and re.match(r"^\s*enable-hwcodec\s*=", line):
            return bool(re.match(r"^\s*enable-hwcodec\s*=\s*['\"]N['\"]\s*(?:#.*)?$", line))
    return False


def update_config(path: Path, expected_uid: int, expected_gid: int) -> str:
    if not path.exists():
        return "absent"
    if path.is_symlink():
        return "unsafe-symlink"
    try:
        info = path.stat()
        if not stat.S_ISREG(info.st_mode) or info.st_uid != expected_uid:
            return "unsafe-owner-or-type"
        original = path.read_text(encoding="utf-8")
    except (OSError, UnicodeError):
        return "unreadable"

    lines = original.splitlines(keepends=True)
    newline = "\r\n" if "\r\n" in original else "\n"
    section_start = None
    section_end = len(lines)
    option_line = None
    for index, line in enumerate(lines):
        match = re.match(r"^\s*\[([^]]+)]\s*(?:\r?\n)?$", line)
        if match:
            if section_start is not None and section_end == len(lines):
                section_end = index
            if match.group(1) == "options":
                section_start = index
                section_end = len(lines)
            elif section_start is not None:
                section_end = index
        elif section_start is not None and section_end == len(lines) and re.match(r"^\s*enable-hwcodec\s*=", line):
            option_line = index

    desired = f"{SETTING} = 'N'{newline}"
    if option_line is not None:
        if lines[option_line] == desired:
            return "already-safe"
        lines[option_line] = desired
    elif section_start is not None:
        lines.insert(section_end, desired)
    else:
        if lines and not lines[-1].endswith(("\n", "\r")):
            lines[-1] += newline
        if lines and lines[-1].strip():
            lines.append(newline)
        lines.extend((f"[options]{newline}", desired))

    path.parent.mkdir(mode=0o700, parents=True, exist_ok=True)
    fd, temporary_name = tempfile.mkstemp(prefix=".RustDesk2.toml.eczos-", dir=path.parent)
    temporary = Path(temporary_name)
    try:
        os.fchmod(fd, stat.S_IMODE(info.st_mode))
        os.fchown(fd, expected_uid, expected_gid)
        with os.fdopen(fd, "w", encoding="utf-8", newline="") as stream:
            stream.writelines(lines)
            stream.flush()
            os.fsync(stream.fileno())
        os.replace(temporary, path)
    finally:
        temporary.unlink(missing_ok=True)
    return "changed"


def service_state() -> tuple[bool, bool]:
    enabled = subprocess.run(
        ["systemctl", "is-enabled", "rustdesk.service"], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL
    ).returncode == 0
    active = subprocess.run(
        ["systemctl", "is-active", "rustdesk.service"], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL
    ).returncode == 0
    return enabled, active


def recent_gpu_crash() -> bool:
    try:
        log = subprocess.run(
            ["journalctl", "-k", "-b", "--no-pager", "-o", "cat"],
            check=False,
            capture_output=True,
            text=True,
            timeout=15,
        ).stdout.lower()
    except (OSError, subprocess.TimeoutExpired):
        return False
    return "ring vce" in log and "timeout" in log and "process rustdesk" in log


def status_payload() -> dict[str, object]:
    installed = rustdesk_installed()
    configured = [str(path) for path, _, _ in candidate_configs() if software_codec_configured(path)]
    enabled, active = service_state() if installed else (False, False)
    return {
        "schemaVersion": 1,
        "installed": installed,
        "serviceEnabled": enabled,
        "serviceActive": active,
        "softwareEncoding": bool(configured),
        "configuredProfiles": len(configured),
        "gpuCrashDetectedThisBoot": recent_gpu_crash(),
    }


def write_state(payload: dict[str, object]) -> None:
    STATE_DIR.mkdir(mode=0o755, parents=True, exist_ok=True)
    payload = dict(payload)
    payload["lastApplied"] = datetime.now(timezone.utc).isoformat()
    temporary = STATE_FILE.with_suffix(".tmp")
    temporary.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
    os.chmod(temporary, 0o644)
    os.replace(temporary, STATE_FILE)


def apply(restart: bool) -> int:
    if os.geteuid() != 0:
        print("Administrative privileges are required.", file=sys.stderr)
        return 2
    results: dict[str, str] = {}
    changed = False
    for path, uid, gid in candidate_configs():
        result = update_config(path, uid, gid)
        results[str(path)] = result
        changed |= result == "changed"
    payload = status_payload()
    payload["profiles"] = results
    write_state(payload)
    if restart and changed and payload["serviceActive"]:
        subprocess.run(["systemctl", "try-restart", "rustdesk.service"], check=False)
    print(json.dumps(status_payload(), separators=(",", ":")))
    return 0


def main() -> int:
    parser = argparse.ArgumentParser(description="ECZOS remote-support stability guard")
    subparsers = parser.add_subparsers(dest="command", required=True)
    status = subparsers.add_parser("status")
    status.add_argument("--json", action="store_true")
    apply_parser = subparsers.add_parser("apply")
    apply_parser.add_argument("--restart", action="store_true")
    arguments = parser.parse_args()
    if arguments.command == "apply":
        return apply(arguments.restart)
    payload = status_payload()
    if arguments.json:
        print(json.dumps(payload, separators=(",", ":")))
    else:
        print("ECZHOA/RustDesk: " + ("installed" if payload["installed"] else "not installed"))
        print("Encoding: " + ("safe software mode" if payload["softwareEncoding"] else "not configured"))
        print("Service: " + ("active" if payload["serviceActive"] else "inactive"))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
