#!/usr/bin/env python3
"""Secure terminal relay for brt -access serve/connect."""

from __future__ import annotations

import argparse
import hmac
import json
import os
import re
import secrets
import shutil
import signal
import socket
import string
import struct
import sys
import threading
import time
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from typing import Any

IS_WINDOWS = sys.platform == "win32"

import select

if IS_WINDOWS:
    import subprocess
else:
    import pty
    import termios
    import tty

MSG_AUTH = 0x01
MSG_AUTH_OK = 0x02
MSG_AUTH_FAIL = 0x03
MSG_DATA = 0x04
MSG_CLOSE = 0x05

DEFAULT_AUDIT_LOG = str(Path.home() / ".brt" / "access.audit.log")
DEFAULT_CLIENT_INFO = str(Path.home() / ".brt" / "access.client")
DEFAULT_ADMIN_SOCK = (
    "tcp:127.0.0.1:8767"
    if IS_WINDOWS
    else str(Path.home() / ".brt" / "access.admin.sock")
)
SESSION_ID_ALPHABET = string.ascii_letters + string.digits
MAX_COMMAND_LOG_LEN = 4096

# ANSI colors (bettercap-style dashboard)
C_RESET = "\033[0m"
C_DIM = "\033[2m"
C_BOLD = "\033[1m"
C_CYAN = "\033[36m"
C_GREEN = "\033[32m"
C_YELLOW = "\033[33m"
C_RED = "\033[31m"
C_MAGENTA = "\033[35m"
C_WHITE = "\033[97m"
C_BG = "\033[40m"


def send_msg(sock: socket.socket, msg_type: int, data: bytes = b"") -> None:
    sock.sendall(struct.pack("!BH", msg_type, len(data)) + data)


def recv_exact(sock: socket.socket, n: int) -> bytes:
    buf = b""
    while len(buf) < n:
        chunk = sock.recv(n - len(buf))
        if not chunk:
            raise ConnectionError("Connection closed")
        buf += chunk
    return buf


def recv_msg(sock: socket.socket) -> tuple[int, bytes]:
    header = recv_exact(sock, 3)
    msg_type, length = struct.unpack("!BH", header)
    data = recv_exact(sock, length) if length else b""
    return msg_type, data


def utc_now_iso() -> str:
    dt = datetime.now(timezone.utc)
    return dt.strftime("%Y-%m-%dT%H:%M:%S.") + f"{dt.microsecond // 1000:03d}Z"


def verify_token(provided: str, expected: str) -> bool:
    return hmac.compare_digest(provided.encode(), expected.encode())


def default_audit_log_path() -> str:
    return os.environ.get("BRT_ACCESS_AUDIT_LOG", DEFAULT_AUDIT_LOG)


def default_client_info_path() -> str:
    return os.environ.get("BRT_ACCESS_CLIENT_INFO", DEFAULT_CLIENT_INFO)


def write_client_info(
    endpoint: str,
    token: str,
    pid: int,
    session_id: str = "",
) -> None:
    path = Path(default_client_info_path())
    path.parent.mkdir(parents=True, exist_ok=True)
    payload = {
        "endpoint": endpoint,
        "token": token,
        "pid": pid,
        "session_id": session_id,
        "started": utc_now_iso(),
    }
    path.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")


def read_client_info() -> dict[str, Any] | None:
    path = Path(default_client_info_path())
    if not path.is_file():
        return None
    try:
        data = json.loads(path.read_text(encoding="utf-8"))
    except (json.JSONDecodeError, OSError):
        return None
    return data if isinstance(data, dict) else None


def clear_client_info() -> None:
    path = Path(default_client_info_path())
    try:
        path.unlink(missing_ok=True)
    except OSError:
        pass


def _pid_alive(pid: int) -> bool:
    if pid <= 0:
        return False
    try:
        os.kill(pid, 0)
        return True
    except OSError:
        return False


def leave_client() -> int:
    info = read_client_info()
    if not info:
        print("error: not connected to any access server", file=sys.stderr)
        print(
            "hint: type exit in the remote shell, or connect first with brt -access connect",
            file=sys.stderr,
        )
        return 1

    pid = int(info.get("pid") or 0)
    session_id = str(info.get("session_id") or "")
    endpoint = str(info.get("endpoint") or "")

    if not _pid_alive(pid):
        clear_client_info()
        print("Disconnected (session was already closed)")
        return 0

    try:
        os.kill(pid, signal.SIGTERM)
    except OSError as exc:
        clear_client_info()
        print(f"error: could not disconnect: {exc}", file=sys.stderr)
        return 1

    for _ in range(30):
        if not _pid_alive(pid):
            break
        time.sleep(0.1)

    if _pid_alive(pid):
        try:
            os.kill(pid, getattr(signal, "SIGKILL", signal.SIGTERM))
        except OSError:
            pass

    clear_client_info()
    sid_msg = f" (session {session_id})" if session_id else ""
    print(f"Disconnected from {endpoint}{sid_msg}")
    return 0


def clear_serve_logs(audit_log: str | None = None, *, quiet: bool = False) -> int:
    audit = audit_log or default_audit_log_path()
    state_dir = Path.home() / ".brt"
    cleared: list[str] = []
    for path in (Path(audit), state_dir / "access.stdout", state_dir / "access.log"):
        path.parent.mkdir(parents=True, exist_ok=True)
        path.write_text("", encoding="utf-8")
        cleared.append(str(path))
    if not quiet:
        print(f"Cleared {len(cleared)} log file(s)", flush=True)
    return 0


def default_admin_sock_path() -> str:
    return os.environ.get("BRT_ACCESS_ADMIN_SOCK", DEFAULT_ADMIN_SOCK)


def _default_shell() -> str:
    if IS_WINDOWS:
        for candidate in (
            os.environ.get("BRT_SHELL"),
            shutil.which("pwsh"),
            shutil.which("powershell"),
            os.environ.get("COMSPEC"),
        ):
            if candidate:
                return candidate
        return "cmd.exe"
    return os.environ.get("SHELL", "/bin/bash")


def _admin_is_tcp(spec: str) -> bool:
    return spec.startswith("tcp:")


def _admin_open_connection(spec: str) -> socket.socket:
    if _admin_is_tcp(spec):
        host, port_str = spec[4:].rsplit(":", 1)
        sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        sock.settimeout(2.0)
        sock.connect((host, int(port_str)))
        return sock
    path = Path(spec)
    if not path.exists():
        raise FileNotFoundError(spec)
    if not hasattr(socket, "AF_UNIX"):
        raise OSError("Unix admin socket not supported on this platform")
    sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
    sock.settimeout(2.0)
    sock.connect(str(path))
    return sock


def _admin_bind_server(spec: str) -> socket.socket:
    if _admin_is_tcp(spec):
        host, port_str = spec[4:].rsplit(":", 1)
        server = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        server.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
        server.bind((host, int(port_str)))
        server.listen(8)
        return server
    path = Path(spec)
    path.parent.mkdir(parents=True, exist_ok=True)
    if path.exists():
        path.unlink()
    if not hasattr(socket, "AF_UNIX"):
        raise OSError("Unix admin socket not supported on this platform")
    server = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
    server.bind(str(path))
    server.listen(8)
    return server


def _admin_close_server(spec: str, server: socket.socket) -> None:
    try:
        server.close()
    except OSError:
        pass
    if not _admin_is_tcp(spec):
        path = Path(spec)
        if path.exists():
            path.unlink()


def new_session_id(taken: set[str]) -> str:
    for _ in range(200):
        sid = "".join(secrets.choice(SESSION_ID_ALPHABET) for _ in range(6))
        if sid not in taken:
            return sid
    raise RuntimeError("could not allocate session id")


class AccessSession:
    def __init__(self, session_id: str, remote: str) -> None:
        self.session_id = session_id
        self.remote = remote
        self.started = time.perf_counter()
        self.last_action = self.started
        self.bytes_in = 0
        self.bytes_out = 0
        self.command_buffer = ""
        self.command_count = 0
        self.shell_pid: int | None = None
        self.shell = _default_shell()


@dataclass
class ActiveSessionHandle:
    session_id: str
    remote: str
    started: float
    shell_pid: int | None = None
    conn: socket.socket | None = None
    master_fd: int | None = None
    proc: Any = None
    stop: threading.Event = field(default_factory=threading.Event)


class SessionRegistry:
    def __init__(self) -> None:
        self._lock = threading.Lock()
        self._sessions: dict[str, ActiveSessionHandle] = {}

    def ids(self) -> set[str]:
        with self._lock:
            return set(self._sessions)

    def register(self, handle: ActiveSessionHandle) -> None:
        with self._lock:
            self._sessions[handle.session_id] = handle

    def unregister(self, session_id: str) -> None:
        with self._lock:
            self._sessions.pop(session_id, None)

    def get(self, session_id: str) -> ActiveSessionHandle | None:
        with self._lock:
            return self._sessions.get(session_id)

    def list_clients(self) -> list[dict[str, Any]]:
        now = time.perf_counter()
        with self._lock:
            out = []
            for sid, handle in sorted(self._sessions.items()):
                uptime_ms = int((now - handle.started) * 1000)
                out.append(
                    {
                        "id": sid,
                        "remote": handle.remote,
                        "shell_pid": handle.shell_pid,
                        "uptime_ms": uptime_ms,
                    }
                )
            return out

    def disconnect(self, session_id: str, audit: AccessAuditLog | None = None) -> bool:
        handle = self.get(session_id)
        if not handle:
            return False
        handle.stop.set()
        if handle.conn is not None:
            try:
                send_msg(handle.conn, MSG_CLOSE)
                handle.conn.shutdown(socket.SHUT_RDWR)
            except OSError:
                pass
        if handle.shell_pid:
            try:
                os.kill(handle.shell_pid, signal.SIGTERM)
            except ProcessLookupError:
                pass
        if handle.proc is not None:
            try:
                handle.proc.terminate()
            except OSError:
                pass
        if audit is not None:
            audit.log(
                "session_disconnect",
                level="warn",
                session=session_id,
                remote=handle.remote,
                reason="admin_disconnect",
            )
        return True


class AccessAuditLog:
    def __init__(self, path: str) -> None:
        self.path = Path(path)
        self.path.parent.mkdir(parents=True, exist_ok=True)
        self.server_started = time.perf_counter()
        self._lock = threading.Lock()

    def _write(self, record: dict) -> None:
        line = json.dumps(record, ensure_ascii=False)
        with self._lock:
            with open(self.path, "a", encoding="utf-8") as fh:
                fh.write(line + "\n")

    def log(self, event: str, level: str = "info", **fields) -> None:
        record = {"ts": utc_now_iso(), "level": level, "event": event}
        record.update(fields)
        self._write(record)

    def session_log(
        self,
        session: AccessSession,
        event: str,
        level: str = "info",
        **fields,
    ) -> None:
        now = time.perf_counter()
        delay_ms = int((now - session.last_action) * 1000)
        since_session_ms = int((now - session.started) * 1000)
        session.last_action = now
        record = {
            "ts": utc_now_iso(),
            "level": level,
            "event": event,
            "session": session.session_id,
            "remote": session.remote,
            "delay_ms": delay_ms,
            "since_session_ms": since_session_ms,
        }
        record.update(fields)
        self._write(record)

    def server_start(self, host: str, port: int, pid: int, admin_sock: str) -> None:
        self.log(
            "server_start",
            host=host,
            port=port,
            pid=pid,
            audit_log=str(self.path),
            admin_sock=admin_sock,
        )

    def server_stop(self, reason: str) -> None:
        uptime_ms = int((time.perf_counter() - self.server_started) * 1000)
        self.log("server_stop", reason=reason, uptime_ms=uptime_ms)


def _sanitize_command(text: str) -> str:
    text = text.replace("\r", "").replace("\n", " ").strip()
    text = re.sub(r"[\x00-\x08\x0b-\x1f\x7f]", "", text)
    if len(text) > MAX_COMMAND_LOG_LEN:
        text = text[:MAX_COMMAND_LOG_LEN] + "…"
    return text


def _log_client_input(session: AccessSession, audit: AccessAuditLog, data: bytes) -> None:
    session.bytes_in += len(data)
    chunk = data.decode("utf-8", errors="replace")
    for ch in chunk:
        if ch in ("\n", "\r"):
            line = _sanitize_command(session.command_buffer)
            session.command_buffer = ""
            if not line:
                continue
            session.command_count += 1
            audit.session_log(
                session,
                "user_command",
                action="command",
                command=line,
                command_index=session.command_count,
            )
        elif ch in ("\x7f", "\b"):
            session.command_buffer = session.command_buffer[:-1]
        elif ch == "\x03":
            audit.session_log(session, "user_input", action="interrupt", key="^C")
        elif ch == "\x1a":
            audit.session_log(session, "user_input", action="suspend", key="^Z")
        elif ch == "\x04":
            audit.session_log(session, "user_input", action="eof", key="^D")
        elif ch.isprintable() or ch == "\t":
            session.command_buffer += ch


def set_raw(enabled: bool):
    if IS_WINDOWS or not sys.stdin.isatty():
        return None
    fd = sys.stdin.fileno()
    old = termios.tcgetattr(fd)
    if enabled:
        tty.setraw(fd)
    else:
        termios.tcsetattr(fd, termios.TCSADRAIN, old)
    return old


def relay_io_windows(sock: socket.socket) -> None:
    stop = threading.Event()

    def _reader() -> None:
        while not stop.is_set():
            try:
                msg_type, data = recv_msg(sock)
                if msg_type == MSG_CLOSE:
                    break
                if msg_type == MSG_DATA:
                    sys.stdout.buffer.write(data)
                    sys.stdout.buffer.flush()
            except (ConnectionError, OSError):
                break

    thread = threading.Thread(target=_reader, daemon=True)
    thread.start()
    try:
        while True:
            data = sys.stdin.buffer.read(1)
            if not data:
                break
            send_msg(sock, MSG_DATA, data)
    finally:
        stop.set()
        thread.join(timeout=0.5)


def relay_io(sock: socket.socket, master_fd=None) -> None:
    if IS_WINDOWS and master_fd is None:
        relay_io_windows(sock)
        return
    if IS_WINDOWS:
        raise RuntimeError("Windows relay_io requires pipe mode")
    old_attrs = set_raw(True)
    try:
        while True:
            rlist = [sock]
            if master_fd is not None:
                rlist.append(master_fd)
            else:
                rlist.append(sys.stdin.fileno())
            readable, _, _ = select.select(rlist, [], [], 0.5)
            if sock in readable:
                msg_type, data = recv_msg(sock)
                if msg_type == MSG_CLOSE:
                    break
                if msg_type == MSG_DATA:
                    os.write(sys.stdout.fileno(), data)
            if master_fd is not None and master_fd in readable:
                try:
                    data = os.read(master_fd, 4096)
                    if not data:
                        break
                    send_msg(sock, MSG_DATA, data)
                except OSError:
                    break
            if master_fd is None and sys.stdin.fileno() in readable:
                data = os.read(sys.stdin.fileno(), 4096)
                if not data:
                    break
                send_msg(sock, MSG_DATA, data)
    finally:
        if old_attrs:
            termios.tcsetattr(sys.stdin.fileno(), termios.TCSADRAIN, old_attrs)


def relay_io_server_process(
    sock: socket.socket,
    proc: subprocess.Popen,
    session: AccessSession,
    audit: AccessAuditLog,
    stop: threading.Event,
) -> str:
    reason = "disconnect"
    write_lock = threading.Lock()

    def _stdout_reader() -> None:
        nonlocal reason
        assert proc.stdout is not None
        while not stop.is_set():
            try:
                data = proc.stdout.read(4096)
                if not data:
                    reason = "shell_exit"
                    stop.set()
                    break
                session.bytes_out += len(data)
                send_msg(sock, MSG_DATA, data)
            except OSError as exc:
                audit.session_log(
                    session,
                    "io_error",
                    level="error",
                    direction="stdout_read",
                    error=str(exc),
                )
                reason = "io_error"
                stop.set()
                break

    reader = threading.Thread(target=_stdout_reader, daemon=True)
    reader.start()
    try:
        while not stop.is_set():
            try:
                msg_type, data = recv_msg(sock)
            except (ConnectionError, OSError) as exc:
                audit.session_log(
                    session,
                    "io_error",
                    level="error",
                    direction="socket",
                    error=str(exc),
                )
                reason = "io_error"
                break
            if msg_type == MSG_CLOSE:
                reason = "client_close"
                break
            if msg_type == MSG_DATA:
                _log_client_input(session, audit, data)
                if proc.stdin is not None:
                    with write_lock:
                        proc.stdin.write(data)
                        proc.stdin.flush()
    finally:
        stop.set()
        reader.join(timeout=0.5)
    if stop.is_set() and reason == "disconnect":
        reason = "admin_disconnect"
    return reason


def relay_io_server(
    sock: socket.socket,
    master_fd: int,
    session: AccessSession,
    audit: AccessAuditLog,
    stop: threading.Event,
) -> str:
    if IS_WINDOWS:
        raise RuntimeError("relay_io_server is Unix-only")
    reason = "disconnect"
    try:
        while not stop.is_set():
            readable, _, _ = select.select([sock, master_fd], [], [], 0.25)
            if stop.is_set():
                reason = "admin_disconnect"
                break
            if sock in readable:
                msg_type, data = recv_msg(sock)
                if msg_type == MSG_CLOSE:
                    reason = "client_close"
                    break
                if msg_type == MSG_DATA:
                    _log_client_input(session, audit, data)
                    os.write(master_fd, data)
            if master_fd in readable:
                try:
                    data = os.read(master_fd, 4096)
                    if not data:
                        reason = "shell_exit"
                        break
                    session.bytes_out += len(data)
                    send_msg(sock, MSG_DATA, data)
                except OSError as exc:
                    audit.session_log(
                        session,
                        "io_error",
                        level="error",
                        direction="pty_read",
                        error=str(exc),
                    )
                    reason = "io_error"
                    break
    except (ConnectionError, OSError) as exc:
        audit.session_log(
            session,
            "io_error",
            level="error",
            direction="socket",
            error=str(exc),
        )
        reason = "io_error"
    if stop.is_set() and reason == "disconnect":
        reason = "admin_disconnect"
    return reason


def handle_client(
    conn: socket.socket,
    addr,
    token: str,
    audit: AccessAuditLog,
    registry: SessionRegistry,
) -> None:
    remote = f"{addr[0]}:{addr[1]}"
    session_id = new_session_id(registry.ids())
    session = AccessSession(session_id, remote)
    handle = ActiveSessionHandle(session_id, remote, session.started, conn=conn)
    connect_started = time.perf_counter()

    audit.session_log(session, "client_connect", action="connect")
    print(f"[+] Connection from {remote} (session {session_id})", file=sys.stderr)

    pid = 0
    master_fd: int | None = None
    proc: subprocess.Popen | None = None
    try:
        auth_started = time.perf_counter()
        msg_type, data = recv_msg(conn)
        if msg_type != MSG_AUTH:
            send_msg(conn, MSG_AUTH_FAIL, b"Expected auth")
            audit.session_log(
                session,
                "auth_fail",
                level="warn",
                reason="expected_auth",
                auth_delay_ms=int((time.perf_counter() - auth_started) * 1000),
            )
            return

        provided = data.decode("utf-8", errors="replace")
        auth_delay_ms = int((time.perf_counter() - auth_started) * 1000)
        if not verify_token(provided, token):
            send_msg(conn, MSG_AUTH_FAIL, b"Invalid token")
            audit.session_log(
                session,
                "auth_fail",
                level="warn",
                reason="invalid_token",
                auth_delay_ms=auth_delay_ms,
            )
            print(f"[!] Auth failed from {remote}", file=sys.stderr)
            return

        send_msg(
            conn,
            MSG_AUTH_OK,
            json.dumps(
                {"message": "Authenticated", "session_id": session_id},
            ).encode(),
        )
        audit.session_log(
            session,
            "auth_ok",
            auth_delay_ms=auth_delay_ms,
            connect_delay_ms=int((time.perf_counter() - connect_started) * 1000),
        )
        print(f"[+] Authenticated {remote} (session {session_id})", file=sys.stderr)

        if IS_WINDOWS:
            shell_cmd = [session.shell]
            if session.shell.lower().endswith(("cmd.exe", "cmd")):
                shell_cmd = [session.shell]
            elif "powershell" in session.shell.lower() or session.shell.lower().endswith("pwsh"):
                shell_cmd = [session.shell, "-NoLogo"]
            proc = subprocess.Popen(
                shell_cmd,
                stdin=subprocess.PIPE,
                stdout=subprocess.PIPE,
                stderr=subprocess.STDOUT,
                text=False,
                bufsize=0,
                creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0),
            )
            session.shell_pid = proc.pid
            handle.shell_pid = proc.pid
            handle.proc = proc
            registry.register(handle)
            audit.session_log(
                session,
                "shell_spawn",
                shell=session.shell,
                shell_pid=proc.pid,
            )
            reason = relay_io_server_process(conn, proc, session, audit, handle.stop)
            exit_code = proc.poll()
            pid = proc.pid
        else:
            pid, master_fd = pty.fork()
            if pid == 0:
                os.execvp(session.shell, [session.shell, "-l"])
                sys.exit(1)

            session.shell_pid = pid
            handle.shell_pid = pid
            handle.master_fd = master_fd
            registry.register(handle)

            audit.session_log(session, "shell_spawn", shell=session.shell, shell_pid=pid)

            reason = relay_io_server(conn, master_fd, session, audit, handle.stop)

            exit_code = None
            try:
                os.kill(pid, signal.SIGTERM)
                _, status = os.waitpid(pid, 0)
                if os.WIFEXITED(status):
                    exit_code = os.WEXITSTATUS(status)
                elif os.WIFSIGNALED(status):
                    exit_code = -os.WTERMSIG(status)
            except (ProcessLookupError, ChildProcessError):
                pass

        duration_ms = int((time.perf_counter() - session.started) * 1000)
        audit.session_log(
            session,
            "session_end",
            reason=reason,
            duration_ms=duration_ms,
            bytes_in=session.bytes_in,
            bytes_out=session.bytes_out,
            commands=session.command_count,
            shell_pid=pid,
            exit_code=exit_code,
        )
        print(
            f"[-] Session {session_id} ended ({reason}, {duration_ms}ms)",
            file=sys.stderr,
        )
    except (ConnectionError, OSError) as exc:
        audit.session_log(session, "session_error", level="error", error=str(exc))
        print(f"[!] Connection error ({remote}): {exc}", file=sys.stderr)
    finally:
        registry.unregister(session_id)
        if proc is not None:
            try:
                if proc.poll() is None:
                    proc.terminate()
            except OSError:
                pass
        if master_fd is not None:
            try:
                os.close(master_fd)
            except OSError:
                pass
        conn.close()


def _admin_handle_connection(conn: socket.socket, registry: SessionRegistry, audit: AccessAuditLog) -> None:
    with conn:
        try:
            raw = conn.recv(65536).decode("utf-8", errors="replace").strip()
            req = json.loads(raw) if raw else {}
            cmd = req.get("cmd", "")
            if cmd == "list":
                resp = {"ok": True, "clients": registry.list_clients()}
            elif cmd == "disconnect":
                sid = str(req.get("id", ""))
                ok = registry.disconnect(sid, audit)
                resp = (
                    {"ok": True, "id": sid}
                    if ok
                    else {"ok": False, "error": f"unknown session {sid}"}
                )
            else:
                resp = {"ok": False, "error": f"unknown cmd {cmd}"}
            conn.sendall((json.dumps(resp) + "\n").encode())
        except (json.JSONDecodeError, OSError):
            conn.sendall(b'{"ok":false,"error":"bad request"}\n')


def admin_server_loop(
    admin_spec: str,
    registry: SessionRegistry,
    audit: AccessAuditLog,
    stop: threading.Event,
) -> None:
    server = _admin_bind_server(admin_spec)
    server.settimeout(0.5)

    while not stop.is_set():
        try:
            conn, _ = server.accept()
        except socket.timeout:
            continue
        except OSError:
            break
        threading.Thread(
            target=_admin_handle_connection,
            args=(conn, registry, audit),
            daemon=True,
        ).start()

    _admin_close_server(admin_spec, server)


def admin_request(admin_spec: str, payload: dict) -> dict:
    try:
        sock = _admin_open_connection(admin_spec)
    except FileNotFoundError:
        return {"ok": False, "error": "access server not running"}
    except OSError as exc:
        return {"ok": False, "error": str(exc)}
    try:
        with sock:
            sock.sendall((json.dumps(payload) + "\n").encode())
            data = sock.recv(65536).decode("utf-8", errors="replace").strip()
        return json.loads(data) if data else {"ok": False, "error": "empty response"}
    except (OSError, json.JSONDecodeError) as exc:
        return {"ok": False, "error": str(exc)}


def _fmt_bytes(n: int) -> str:
    if n < 1024:
        return f"{n}B"
    if n < 1024 * 1024:
        return f"{n // 1024}K"
    return f"{n / (1024 * 1024):.1f}M"


def _fmt_uptime(ms: int) -> str:
    sec = max(0, ms // 1000)
    if sec < 60:
        return f"{sec}s"
    mins, sec = divmod(sec, 60)
    if mins < 60:
        return f"{mins}m{sec:02d}s"
    hrs, mins = divmod(mins, 60)
    return f"{hrs}h{mins:02d}m"


def _event_detail(rec: dict) -> str:
    event = rec.get("event", "")
    if event == "server_start":
        host = rec.get("host", "")
        port = rec.get("port", "")
        if host != "" and port != "":
            return f"{host}:{port}"
    if event == "user_command":
        return rec.get("command", "")
    if event in ("auth_fail", "session_end", "session_disconnect"):
        return rec.get("reason", "") or rec.get("error", "")
    if event == "user_input":
        return rec.get("key", rec.get("action", ""))
    for key in ("command", "error", "shell", "host", "reason"):
        if rec.get(key):
            return str(rec[key])
    return ""


def _level_color(level: str) -> str:
    if level == "error":
        return C_RED
    if level == "warn":
        return C_YELLOW
    if level == "info":
        return C_GREEN
    return C_WHITE


def _event_color(event: str) -> str:
    if event in ("user_command", "user_input"):
        return C_CYAN
    if event in ("auth_fail", "session_error", "io_error"):
        return C_RED
    if event in ("session_end", "session_disconnect"):
        return C_YELLOW
    if event in ("auth_ok", "client_connect", "shell_spawn"):
        return C_GREEN
    return C_WHITE


def _parse_audit_line(line: str) -> dict | None:
    line = line.strip()
    if not line:
        return None
    try:
        return json.loads(line)
    except json.JSONDecodeError:
        return None


def _short_ts(ts: str) -> str:
    if "T" in ts:
        return ts.split("T", 1)[1].replace("Z", "")[:12]
    return ts[:12]


def _draw_dashboard(
    cols: int,
    rows: int,
    clients: list[dict],
    events: list[dict],
    audit_path: str,
    server_online: bool,
    *,
    endpoint: str = "",
    token: str = "",
    status_msg: str = "",
    cmd_input: str = "",
    interactive: bool = False,
) -> None:
    out: list[str] = []
    title = f"{C_BOLD}{C_CYAN}━━ brt access — live audit ━━{C_RESET}"
    if endpoint:
        title += f"  {C_GREEN}{endpoint}{C_RESET}"
    out.append(title)
    if token:
        tok_show = token if len(token) <= 20 else f"{token[:8]}…{token[-4:]}"
        out.append(
            f"{C_DIM}token {C_WHITE}{tok_show}{C_RESET}  ·  "
            f"{audit_path}  ·  "
            f"{'online' if server_online else 'offline'}{C_RESET}"
        )
    else:
        out.append(
            f"{C_DIM}{audit_path}  ·  "
            f"{'server online' if server_online else 'server offline'}"
            f"{'  ·  q quit' if not interactive else ''}{C_RESET}"
        )

    client_line = f"{C_BOLD} Connected ({len(clients)}){C_RESET} "
    if clients:
        parts = []
        for c in clients[:8]:
            parts.append(f"{C_MAGENTA}{c['id']}{C_RESET}@{c['remote']}")
        client_line += "  ".join(parts)
        if len(clients) > 8:
            client_line += f"  {C_DIM}+{len(clients) - 8} more{C_RESET}"
    else:
        client_line += f"{C_DIM}(none){C_RESET}"
    out.append(client_line)
    out.append("")

    out.append(f"{C_BOLD} Events{C_RESET}")
    out.append(
        f"{C_DIM}  {'TIME':<12}  {'ID':<6}  {'EVENT':<16}  {'DETAIL':<28}  +ms{C_RESET}"
    )

    footer_lines = 2 if interactive else 0
    view_h = max(3, rows - len(out) - footer_lines - 1)
    for rec in events[-view_h:]:
        ts = _short_ts(str(rec.get("ts", "")))
        sid = str(rec.get("session", ""))[:6]
        event = str(rec.get("event", ""))[:16]
        detail = _event_detail(rec)[: max(8, cols - 58)]
        delay = rec.get("delay_ms", "")
        delay_s = str(delay) if delay != "" else ""
        lc = _level_color(str(rec.get("level", "info")))
        ec = _event_color(event)
        out.append(
            f"  {C_DIM}{ts:<12}{C_RESET}  "
            f"{C_MAGENTA}{sid:<6}{C_RESET}  "
            f"{ec}{event:<16}{C_RESET}  "
            f"{lc}{detail:<28}{C_RESET}  "
            f"{C_DIM}{delay_s}{C_RESET}"
        )

    body_lines = rows - footer_lines - 1
    while len(out) < body_lines:
        out.append("")

    if interactive:
        hint = "connect · stop · disconnect <id> · clear · help · quit"
        if status_msg:
            out.append(f"{C_YELLOW}{status_msg[: cols - 1]}{C_RESET}")
        else:
            out.append(f"{C_DIM}{hint[: cols - 1]}{C_RESET}")
        prompt = f"{C_BOLD}>{C_RESET} {cmd_input}"
        out.append(prompt)

    screen = "\033[H\033[2J\033[?25l" + "\n".join(out[: rows - 1])
    if interactive:
        screen += "\033[?25h"
        prompt_row = rows
        cursor_col = len(f"> {cmd_input}") + 1
        screen += f"\033[{prompt_row};{cursor_col}H"
    else:
        screen += "\033[?25h"
    sys.stdout.write(screen)
    sys.stdout.flush()


def _copy_to_clipboard(text: str) -> bool:
    import subprocess

    candidates: list[list[str]] = []
    if IS_WINDOWS:
        candidates.append(
            [
                "powershell",
                "-NoProfile",
                "-Command",
                "Set-Clipboard -Value ([Console]::In.ReadToEnd())",
            ]
        )
        if shutil.which("clip"):
            candidates.insert(0, ["clip"])
    if shutil.which("pbcopy"):
        candidates.append(["pbcopy"])
    if shutil.which("wl-copy"):
        candidates.append(["wl-copy"])
    if shutil.which("xclip"):
        candidates.append(["xclip", "-selection", "clipboard"])
    if shutil.which("xsel"):
        candidates.append(["xsel", "--clipboard", "--input"])

    data = text.encode("utf-8")
    for cmd in candidates:
        try:
            if IS_WINDOWS and cmd[0] == "clip":
                subprocess.run(cmd, input=data, check=True, timeout=2)
            elif IS_WINDOWS and "Set-Clipboard" in " ".join(cmd):
                subprocess.run(
                    [
                        "powershell",
                        "-NoProfile",
                        "-Command",
                        f"Set-Clipboard -Value {json.dumps(text)}",
                    ],
                    check=True,
                    timeout=2,
                )
            else:
                subprocess.run(cmd, input=data, check=True, timeout=2)
            return True
        except (OSError, subprocess.SubprocessError):
            continue
    return False


def _connect_command(endpoint: str, token: str) -> str:
    cli = os.environ.get("BRT_CONNECT_CLI")
    if not cli:
        cli = "brtwin" if IS_WINDOWS else "brt"
    return f"{cli} -access connect {endpoint} {token}"


def _copy_connect_command(endpoint: str, token: str) -> tuple[str, bool]:
    if not endpoint or not token:
        return "Connect info not available", False
    cmd = _connect_command(endpoint, token)
    if _copy_to_clipboard(cmd):
        return "Copied connect command to clipboard", True
    return cmd, False


def _stop_server_pid(pid: int) -> bool:
    if pid <= 0:
        return False
    try:
        os.kill(pid, signal.SIGTERM)
    except ProcessLookupError:
        return True
    except OSError:
        return False
    for _ in range(30):
        time.sleep(0.1)
        try:
            os.kill(pid, 0)
        except ProcessLookupError:
            return True
    try:
        os.kill(pid, signal.SIGKILL)
    except ProcessLookupError:
        pass
    return True


def _exec_console_command(
    line: str,
    admin_sock: str,
    server_pid: int | None,
    owns_server: bool,
    endpoint: str = "",
    token: str = "",
    audit_log: str = "",
    console_state: dict[str, bool] | None = None,
) -> tuple[str, int | None]:
    raw = line.strip()
    if not raw:
        return "", None

    parts = raw.split()
    cmd = parts[0].lower()

    if cmd in ("quit", "q", "exit"):
        if owns_server and server_pid:
            _stop_server_pid(server_pid)
            return "Server stopped", 0
        return "", 0

    if cmd == "stop":
        if server_pid:
            _stop_server_pid(server_pid)
            return "Server stopped", 0
        return "No server process", None

    if cmd in ("help", "?"):
        return "connect · stop · disconnect <id> · clear · quit", None

    if cmd in ("clear", "clear-logs", "cls"):
        clear_serve_logs(audit_log or default_audit_log_path(), quiet=True)
        if console_state is not None:
            console_state["clear_logs"] = True
        return "Logs cleared", None

    if cmd == "connect":
        msg, _ok = _copy_connect_command(endpoint, token)
        return msg, None

    session_id = ""
    if cmd in ("disconnect", "dc", "kick"):
        if len(parts) < 2:
            return "Usage: disconnect <6-char-id>", None
        session_id = parts[1]
    elif re.fullmatch(r"[A-Za-z0-9]{6}", parts[0]):
        session_id = parts[0]

    if session_id:
        if not re.fullmatch(r"[A-Za-z0-9]{6}", session_id):
            return "Session id must be 6 alphanumeric characters", None
        resp = admin_request(admin_sock, {"cmd": "disconnect", "id": session_id})
        if resp.get("ok"):
            return f"Disconnected {session_id}", None
        return f"Error: {resp.get('error', 'disconnect failed')}", None

    return f"Unknown command: {raw} (try help)", None


def _read_stdin_keys(timeout: float) -> list[bytes]:
    chunks: list[bytes] = []
    if IS_WINDOWS:
        import msvcrt

        deadline = time.perf_counter() + timeout
        while time.perf_counter() < deadline or chunks:
            if msvcrt.kbhit():
                chunks.append(msvcrt.getch())
                deadline = time.perf_counter() + 0.05
            elif chunks:
                break
            else:
                time.sleep(0.03)
        return chunks
    try:
        fd = sys.stdin.fileno()
    except (ValueError, OSError):
        return chunks
    while True:
        readable, _, _ = select.select([fd], [], [], timeout if not chunks else 0.0)
        if not readable:
            break
        ch = os.read(fd, 64)
        if not ch:
            break
        chunks.append(ch)
        timeout = 0.0
    return chunks


def _apply_input_keys(keys: list[bytes], buf: str) -> tuple[str, bool]:
    """Return (new_buffer, submit)."""
    submit = False
    for chunk in keys:
        for byte in chunk:
            if byte in (10, 13):
                submit = True
            elif byte in (127, 8):
                buf = buf[:-1]
            elif byte == 3:
                if buf:
                    buf = ""
                else:
                    submit = True
            elif 32 <= byte <= 126:
                buf += chr(byte)
    return buf, submit


def _enable_windows_vt() -> None:
    if not IS_WINDOWS:
        return
    try:
        import ctypes

        kernel32 = ctypes.windll.kernel32
        handle = kernel32.GetStdHandle(-11)
        mode = ctypes.c_uint()
        if kernel32.GetConsoleMode(handle, ctypes.byref(mode)):
            kernel32.SetConsoleMode(handle, mode.value | 0x0004)
    except (OSError, AttributeError):
        pass


def watch_logs(
    audit_log: str,
    admin_sock: str,
    follow: bool = True,
    *,
    interactive: bool = False,
    server_pid: int | None = None,
    endpoint: str = "",
    token: str = "",
    owns_server: bool = False,
) -> int:
    path = Path(audit_log)
    if not path.exists():
        path.parent.mkdir(parents=True, exist_ok=True)
        path.touch()

    events: list[dict] = []
    pos = 0
    if path.stat().st_size > 0:
        with open(path, encoding="utf-8") as fh:
            for line in fh:
                rec = _parse_audit_line(line)
                if rec:
                    events.append(rec)
        events = events[-500:]
        pos = path.stat().st_size

    cmd_input = ""
    status_msg = ""
    copied_connect = False
    console_state: dict[str, bool] = {"clear_logs": False}
    old_tty = None
    if interactive and sys.stdin.isatty():
        if IS_WINDOWS:
            _enable_windows_vt()
        else:
            old_tty = termios.tcgetattr(sys.stdin.fileno())
            tty.setcbreak(sys.stdin.fileno())

    exit_code = 0
    try:
        while True:
            cols = shutil.get_terminal_size(fallback=(100, 30)).columns
            rows = shutil.get_terminal_size(fallback=(100, 30)).lines

            if server_pid and server_pid > 0:
                try:
                    os.kill(server_pid, 0)
                except ProcessLookupError:
                    status_msg = "Server exited"
                    _draw_dashboard(
                        cols,
                        rows,
                        [],
                        events,
                        audit_log,
                        False,
                        endpoint=endpoint,
                        token=token,
                        status_msg=status_msg,
                        cmd_input=cmd_input,
                        interactive=interactive,
                    )
                    exit_code = 0
                    break

            admin = admin_request(admin_sock, {"cmd": "list"})
            clients = admin.get("clients", []) if admin.get("ok") else []
            server_online = admin.get("ok", False)

            with open(path, encoding="utf-8") as fh:
                fh.seek(pos)
                for line in fh:
                    rec = _parse_audit_line(line)
                    if rec:
                        events.append(rec)
                pos = fh.tell()
            events = events[-500:]

            if (
                interactive
                and owns_server
                and endpoint
                and token
                and not copied_connect
            ):
                msg, copied_connect = _copy_connect_command(endpoint, token)
                if copied_connect:
                    status_msg = msg

            _draw_dashboard(
                cols,
                rows,
                clients,
                events,
                audit_log,
                server_online,
                endpoint=endpoint,
                token=token,
                status_msg=status_msg,
                cmd_input=cmd_input,
                interactive=interactive,
            )

            if not follow:
                break

            keys = _read_stdin_keys(0.35)
            if interactive:
                if keys and not cmd_input:
                    for chunk in keys:
                        if b"\x03" in chunk:
                            if owns_server and server_pid:
                                _stop_server_pid(server_pid)
                            exit_code = 0
                            follow = False
                            break
                    if not follow:
                        break
                cmd_input, submit = _apply_input_keys(keys, cmd_input)
                if submit:
                    status_msg, done = _exec_console_command(
                        cmd_input,
                        admin_sock,
                        server_pid,
                        owns_server,
                        endpoint=endpoint,
                        token=token,
                        audit_log=audit_log,
                        console_state=console_state,
                    )
                    cmd_input = ""
                    if console_state.get("clear_logs"):
                        events.clear()
                        pos = 0
                        console_state["clear_logs"] = False
                    if done is not None:
                        exit_code = done
                        break
            elif keys:
                for chunk in keys:
                    if b"q" in chunk or b"Q" in chunk or b"\x03" in chunk:
                        exit_code = 0
                        follow = False
                        break
    finally:
        if old_tty:
            termios.tcsetattr(sys.stdin.fileno(), termios.TCSADRAIN, old_tty)
        sys.stdout.write("\033[?25h\n")
        sys.stdout.flush()
    return exit_code


def print_clients(admin_sock: str) -> int:
    resp = admin_request(admin_sock, {"cmd": "list"})
    if not resp.get("ok"):
        print(f"error: {resp.get('error', 'unknown')}", file=sys.stderr)
        return 1
    clients = resp.get("clients", [])
    if not clients:
        print("No connected sessions.")
        return 0
    print(f"{'ID':<6}  {'REMOTE':<24}  {'UPTIME':<8}  PID")
    for c in clients:
        print(
            f"{c['id']:<6}  {c['remote']:<24}  "
            f"{_fmt_uptime(c.get('uptime_ms', 0)):<8}  {c.get('shell_pid') or '-'}"
        )
    return 0


def disconnect_client(admin_sock: str, session_id: str) -> int:
    if not re.fullmatch(r"[A-Za-z0-9]{6}", session_id):
        print("error: session id must be 6 alphanumeric characters", file=sys.stderr)
        return 1
    resp = admin_request(admin_sock, {"cmd": "disconnect", "id": session_id})
    if not resp.get("ok"):
        print(f"error: {resp.get('error', 'disconnect failed')}", file=sys.stderr)
        return 1
    print(f"Disconnected session {session_id}")
    return 0


def serve(host: str, port: int, token: str, audit_log: str, admin_sock: str) -> None:
    audit = AccessAuditLog(audit_log)
    registry = SessionRegistry()
    admin_stop = threading.Event()

    audit.server_start(host, port, os.getpid(), admin_sock)
    print(f"BRT_ACCESS_AUDIT_LOG={audit_log}", flush=True)
    print(f"BRT_ACCESS_ADMIN_SOCK={admin_sock}", flush=True)

    admin_thread = threading.Thread(
        target=admin_server_loop,
        args=(admin_sock, registry, audit, admin_stop),
        daemon=True,
    )
    admin_thread.start()

    server = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
    server.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
    server.bind((host, port))
    server.listen(16)

    def shutdown(sig, frame):
        sig_name = signal.Signals(sig).name if hasattr(signal, "Signals") else str(sig)
        print(f"\n[-] Shutting down ({sig_name})", file=sys.stderr)
        admin_stop.set()
        audit.server_stop(reason=sig_name)
        server.close()
        sys.exit(0)

    signal.signal(signal.SIGINT, shutdown)
    signal.signal(signal.SIGTERM, shutdown)

    while True:
        try:
            conn, addr = server.accept()
            thread = threading.Thread(
                target=handle_client,
                args=(conn, addr, token, audit, registry),
                daemon=True,
            )
            thread.start()
        except OSError:
            admin_stop.set()
            audit.server_stop(reason="socket_error")
            break


def connect(endpoint: str, token: str) -> None:
    if ":" in endpoint:
        host, port_str = endpoint.rsplit(":", 1)
        port = int(port_str)
    else:
        host = endpoint
        port = 8765

    write_client_info(endpoint, token, os.getpid())

    sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)

    def _disconnect_on_term(signum, frame) -> None:
        try:
            send_msg(sock, MSG_CLOSE)
        except OSError:
            pass
        clear_client_info()
        sys.exit(0)

    if hasattr(signal, "SIGTERM"):
        signal.signal(signal.SIGTERM, _disconnect_on_term)

    try:
        sock.connect((host, port))
        send_msg(sock, MSG_AUTH, token.encode())

        msg_type, data = recv_msg(sock)
        if msg_type != MSG_AUTH_OK:
            print(f"Authentication failed: {data.decode()}", file=sys.stderr)
            sys.exit(1)

        session_id = ""
        try:
            payload = json.loads(data.decode("utf-8"))
            if isinstance(payload, dict):
                session_id = str(payload.get("session_id") or "")
        except json.JSONDecodeError:
            pass

        write_client_info(endpoint, token, os.getpid(), session_id)
        hint = f"Connected to {endpoint}"
        if session_id:
            hint += f" (session {session_id})"
        hint += ". Type exit to disconnect, or run `brt -access leave` from another terminal."
        print(hint, file=sys.stderr)

        relay_io(sock)
    finally:
        clear_client_info()
        try:
            send_msg(sock, MSG_CLOSE)
        except OSError:
            pass
        sock.close()


def main() -> int:
    parser = argparse.ArgumentParser()
    sub = parser.add_subparsers(dest="mode", required=True)

    s = sub.add_parser("serve")
    s.add_argument("--host", default="0.0.0.0")
    s.add_argument("--port", type=int, default=8765)
    s.add_argument("--token", default="")
    s.add_argument("--audit-log", default=default_audit_log_path())
    s.add_argument("--admin-sock", default=default_admin_sock_path())

    c = sub.add_parser("connect")
    c.add_argument("endpoint")
    c.add_argument("token")

    w = sub.add_parser("watch-logs")
    w.add_argument("--audit-log", default=default_audit_log_path())
    w.add_argument("--admin-sock", default=default_admin_sock_path())
    w.add_argument("--no-follow", action="store_true")
    w.add_argument("--interactive", action="store_true")
    w.add_argument("--server-pid", type=int, default=0)
    w.add_argument("--endpoint", default="")
    w.add_argument("--token", default="")
    w.add_argument("--owns-server", action="store_true")

    cl = sub.add_parser("clients")
    cl.add_argument("--admin-sock", default=default_admin_sock_path())

    d = sub.add_parser("disconnect")
    d.add_argument("session_id", nargs="?")
    d.add_argument("--admin-sock", default=default_admin_sock_path())

    sub.add_parser("leave")

    lg = sub.add_parser("clear-logs")
    lg.add_argument("--audit-log", default=default_audit_log_path())

    args = parser.parse_args()

    if args.mode == "serve":
        token = args.token or os.urandom(24).hex()
        print(f"BRT_ACCESS_TOKEN={token}", flush=True)
        print(f"BRT_ACCESS_ENDPOINT=0.0.0.0:{args.port}", flush=True)
        serve(args.host, args.port, token, args.audit_log, args.admin_sock)
        return 0
    if args.mode == "connect":
        connect(args.endpoint, args.token)
        return 0
    if args.mode == "watch-logs":
        return watch_logs(
            args.audit_log,
            args.admin_sock,
            follow=not args.no_follow,
            interactive=args.interactive,
            server_pid=args.server_pid or None,
            endpoint=args.endpoint,
            token=args.token,
            owns_server=args.owns_server,
        )
    if args.mode == "clients":
        return print_clients(args.admin_sock)
    if args.mode == "leave":
        return leave_client()
    if args.mode == "disconnect":
        if args.session_id:
            return disconnect_client(args.admin_sock, args.session_id)
        return leave_client()
    if args.mode == "clear-logs":
        return clear_serve_logs(args.audit_log)
    return 1


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