import fcntl
import hashlib
import json
import os
import pty
import re
import select
import signal
import socket
import struct
import subprocess
import sys
import termios
import time
import tty
from collections import deque
from dataclasses import dataclass, field
from pathlib import Path

RELOAD, STOP, RELAUNCH, HEAL = 75, 76, 77, 78
QUICK = 30.0
UNHEALED_LIMIT = 5
UNHEALED_BACKOFF = 2.0
GRACE, STEP = 3.0, 0.05
EXITING = 8.0
CLEAR_LINE = b"\x05\x15"
LONGEST = 65536
INBOX = 1 << 20
TYPED_EVERY = 1.0
HEADLESS_SIZE = (40, 120)
PRINTED, SCREEN, SCREEN_SHAPE, TYPED, LAUNCHED, KEYED = "printed", "screen", "screen.json", "typed", "launched.json", "keyed"
ESCAPES = re.compile(rb"\x1b(?:\[[\x30-\x3f]*[\x20-\x2f]*[\x40-\x7e]|\][^\x07\x1b]*(?:\x07|\x1b\\)|[P_^X][^\x1b]*\x1b\\|O[\x40-\x7e]|[@-_])")


def typing(data: bytes) -> bool:
    return any(byte >= 0x20 for byte in ESCAPES.sub(b"", data))


def waited(pid: int, seconds: float, drain) -> int | None:
    until = time.time() + seconds
    while time.time() < until:
        drain()
        ended, status = os.waitpid(pid, os.WNOHANG)
        if ended:
            return status
        time.sleep(STEP)
    return None


def stop(pid: int, grace: float = GRACE, drain=lambda: None) -> int:
    for sent in (signal.SIGHUP, signal.SIGTERM, signal.SIGKILL):
        try:
            os.kill(pid, sent)
        except ProcessLookupError:
            break
        status = waited(pid, grace, drain)
        if status is not None:
            return status
    try:
        return os.waitpid(pid, 0)[1]
    except ChildProcessError:
        return 0


def as_bytes(value):
    return bytes.fromhex(value) if isinstance(value, str) else value


@dataclass(frozen=True)
class Adopted:
    pid: int
    fd: int
    session: str
    saved: list | None

    @classmethod
    def from_json(cls, raw: dict) -> "Adopted":
        kept = raw["saved"]
        saved = [*kept[:6], [as_bytes(c) for c in kept[6]]] if kept else None
        return cls(raw["pid"], raw["fd"], raw["session"], saved)


@dataclass(frozen=True)
class Start:
    command: list = field(default_factory=list)
    environ: dict = field(default_factory=dict)
    launch: int = 0
    exit: str = ""

    @classmethod
    def from_json(cls, raw: dict) -> "Start":
        return cls(**{key: raw[key] for key in ("command", "environ", "launch", "exit") if key in raw}) if "command" in raw else cls()


@dataclass(frozen=True)
class Launch:
    root: Path
    cwd: Path
    env: str
    agent: str
    worker: list
    heal: list
    ended: list
    args: list
    start: Start
    adopt: Adopted | None
    headless: bool = False

    @classmethod
    def from_json(cls, raw: dict) -> "Launch":
        return cls(Path(raw["root"]), Path(raw["cwd"]), raw["env"], raw["agent"], raw["worker"], raw["heal"], raw["ended"], raw["args"],
                   Start.from_json(raw), Adopted.from_json(raw["adopt"]) if "adopt" in raw else None, bool(raw.get("headless")))


class Supervisor:
    def __init__(self, spec: Launch):
        self.root, self.cwd = spec.root, spec.cwd
        self.env, self.agent = spec.env, spec.agent
        self.worker_command, self.heal_command, self.ended_command = spec.worker, spec.heal, spec.ended
        self.stdin, self.stdout = sys.stdin.fileno(), sys.stdout.fileno()
        self.headless = spec.headless
        adopted = spec.adopt
        self.saved = adopted.saved if adopted else None
        self.pid, self.fd = (adopted.pid, adopted.fd) if adopted else self.spawn(spec.start.command, spec.start.environ)
        self.session = adopted.session if adopted else f"{self.agent}-{self.pid}"
        self.command, self.args, self.launch, self.exit = spec.start.command, spec.args, spec.start.launch, spec.start.exit
        self.worker = None
        self.worker_began = 0.0
        self.unhealed = 0
        self.worker_due = 0.0
        self.alive, self.lifeline = os.pipe()
        os.set_inheritable(self.alive, True)
        self.folder = self.root / "runtime" / "sessions" / self.session
        self.folder.mkdir(parents=True, exist_ok=True)
        self.printed = (self.folder / PRINTED).open("ab")
        self.screen = (self.folder / SCREEN).open("ab")
        self.inbox = self.listen()
        self.sources = [self.fd, self.inbox] if self.headless else [self.fd, self.inbox, self.stdin]
        self.typed_at = 0.0
        self.pending: deque[bytes] = deque()
        self.record_launch()
        self.resize(self.fd)

    def record_launch(self) -> None:
        written = self.folder / f"{LAUNCHED}.writing"
        written.write_text(json.dumps({"pid": self.pid, "command": self.command, "args": self.args, "cwd": str(self.cwd), "launch": self.launch, "env": self.env}))
        os.replace(written, self.folder / LAUNCHED)

    def spawn(self, command: list[str], environ: dict) -> tuple[int, int]:
        pid, fd = pty.fork()
        if pid == 0:
            os.chdir(self.cwd)
            os.execvpe(command[0], command, environ)
        return pid, fd

    def socket_path(self) -> Path:
        shared = Path("/tmp") / f"journal-{hashlib.sha1(str(self.root.resolve()).encode()).hexdigest()[:16]}"
        folder = shared if not shared.exists() or shared.stat().st_uid == os.getuid() else shared.with_name(f"{shared.name}-{os.getuid()}")
        return folder / f"typist-{self.session}.sock"

    def listen(self) -> socket.socket:
        where = self.socket_path()
        where.parent.mkdir(mode=0o700, parents=True, exist_ok=True)
        where.unlink(missing_ok=True)
        inbox = socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM)
        inbox.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, INBOX)
        inbox.bind(str(where))
        inbox.setblocking(False)
        return inbox

    def received(self) -> list[bytes]:
        packets = []
        while True:
            try:
                packets.append(self.inbox.recv(LONGEST))
            except (BlockingIOError, InterruptedError):
                return packets

    def resize(self, fd: int) -> None:
        try:
            size = struct.pack("HHHH", *HEADLESS_SIZE, 0, 0) if self.headless else fcntl.ioctl(self.stdout, termios.TIOCGWINSZ, b"\0" * 8)
            fcntl.ioctl(fd, termios.TIOCSWINSZ, size)
            rows, cols = struct.unpack("HHHH", size)[:2]
        except OSError:
            return
        (self.folder / SCREEN_SHAPE).write_text(json.dumps({"rows": rows, "cols": cols, "at": self.screen.tell(), "printed": self.printed.tell()}))

    def start_worker(self) -> None:
        command = [*self.worker_command, str(self.root), str(self.cwd), self.env, self.agent, self.session, str(self.alive)]
        self.worker = subprocess.Popen(command, cwd=self.cwd, pass_fds=(self.alive,))
        self.worker_began = time.time()

    def worker_ended(self) -> int | None:
        return self.worker.poll() if self.worker else None

    def stop_agent(self) -> int:
        status = self.exited()
        return stop(self.pid, drain=self.drain) if status is None else status

    def exited(self) -> int | None:
        if not self.exit:
            return None
        try:
            os.write(self.fd, CLEAR_LINE)
            time.sleep(STEP)
            os.write(self.fd, self.exit.encode())
            time.sleep(STEP * 6)
            os.write(self.fd, b"\r")
        except OSError:
            return None
        return waited(self.pid, EXITING, self.drain)

    def drain(self) -> None:
        while select.select([self.fd], [], [], 0)[0]:
            try:
                data = os.read(self.fd, LONGEST)
            except OSError:
                return
            if not data:
                return
            self.output(data)

    def relaunch(self) -> None:
        asked = self.folder / "relaunch.json"
        start = Start.from_json(json.loads(asked.read_text()))
        asked.unlink()
        self.stop_agent()
        os.close(self.fd)
        self.pid, self.fd = self.spawn(start.command, start.environ)
        self.command, self.launch, self.exit = start.command, start.launch, start.exit
        self.record_launch()
        self.resize(self.fd)

    def stop_worker(self) -> None:
        if self.worker and self.worker.poll() is None:
            self.worker.terminate()
            self.worker.wait()

    def relay(self) -> int | None:
        ready, writable, _ = select.select(self.sources, [self.fd] if self.pending else [], [], 0.5)
        if self.fd in writable:
            self.feed()
        if self.inbox in ready:
            self.pending.extend(self.received())
            (self.folder / KEYED).touch()
        if self.fd in ready:
            try:
                data = os.read(self.fd, LONGEST)
            except OSError:
                data = b""
            if not data:
                return os.waitpid(self.pid, 0)[1]
            self.output(data)
        if self.stdin in ready:
            data = os.read(self.stdin, LONGEST)
            if not data:
                return self.stop_agent()
            self.pending.append(data)
            (self.folder / KEYED).touch()
            self.typing(data)
        return None

    def feed(self) -> None:
        flags = fcntl.fcntl(self.fd, fcntl.F_GETFL)
        fcntl.fcntl(self.fd, fcntl.F_SETFL, flags | os.O_NONBLOCK)
        try:
            sent = os.write(self.fd, self.pending[0])
        except (BlockingIOError, InterruptedError):
            sent = 0
        except OSError:
            sent = len(self.pending[0])
        finally:
            fcntl.fcntl(self.fd, fcntl.F_SETFL, flags)
        rest = self.pending.popleft()[sent:]
        if rest:
            self.pending.appendleft(rest)

    def output(self, data: bytes) -> None:
        os.write(self.stdout, data)
        self.screen.write(data)
        self.screen.flush()
        self.printed.write(data)
        self.printed.flush()

    def typing(self, data: bytes) -> None:
        typed = self.folder / TYPED
        if b"\r" in data or b"\n" in data:
            typed.unlink(missing_ok=True)
        elif typing(data) and time.time() - self.typed_at >= TYPED_EVERY:
            typed.touch()
            self.typed_at = time.time()

    def delegate(self, command: list[str]) -> str:
        done = subprocess.run(command, cwd=self.cwd, capture_output=True, text=True, timeout=120)
        if done.returncode:
            os.write(self.stdout, f"\r\njournal: {' '.join(command[-2:])} failed: {done.stderr.strip()[-400:]}\r\n".encode())
        return done.stdout.strip()

    def start_ended(self) -> None:
        self.folder.mkdir(parents=True, exist_ok=True)
        try:
            with open(self.folder / "ended.log", "ab") as log:
                subprocess.Popen(self.ended_command, cwd=self.cwd, stdin=subprocess.DEVNULL, stdout=log, stderr=log, start_new_session=True)
        except OSError as failure:
            os.write(self.stdout, f"\r\njournal: {' '.join(self.ended_command[-2:])} did not start: {failure}\r\n".encode())

    def after_worker(self, code: int) -> int | None:
        if code == STOP:
            return self.stop_agent()
        if code == RELAUNCH:
            self.relaunch()
        elif code == HEAL or (code != RELOAD and time.time() - self.worker_began < QUICK):
            line = self.delegate(self.heal_command)
            if line:
                self.unhealed = 0
                os.write(self.stdout, f"\r\n{line}\r\n".encode())
            else:
                self.unhealed += 1
            if self.unhealed >= UNHEALED_LIMIT:
                os.write(self.stdout, b"\r\njournal: this build keeps failing to start and there is no earlier build to go back to, so the session is ending\r\n")
                return self.stop_agent()
        else:
            self.unhealed = 0
        self.worker = None
        self.worker_due = time.time() + (UNHEALED_BACKOFF * 2 ** (self.unhealed - 1) if self.unhealed else 0)
        return None

    def run(self) -> int:
        saved = self.saved
        try:
            saved = saved if saved is not None else termios.tcgetattr(self.stdin)
            tty.setraw(self.stdin)
        except termios.error:
            pass
        signal.signal(signal.SIGWINCH, lambda *_: self.resize(self.fd))
        signal.signal(signal.SIGHUP, lambda *_: sys.exit(129))
        signal.signal(signal.SIGTERM, lambda *_: sys.exit(143))
        status = None
        self.start_worker()
        try:
            while status is None:
                status = self.relay()
                code = self.worker_ended()
                if status is None and code is not None:
                    status = self.after_worker(code)
                if status is None and self.worker is None and time.time() >= self.worker_due:
                    self.start_worker()
        finally:
            self.stop_worker()
            if status is None:
                status = self.stop_agent()
            if saved is not None:
                termios.tcsetattr(self.stdin, termios.TCSADRAIN, saved)
            self.inbox.close()
            self.socket_path().unlink(missing_ok=True)
            self.start_ended()
        return os.waitstatus_to_exitcode(status)


if __name__ == "__main__":
    sys.exit(Supervisor(Launch.from_json(json.loads(sys.argv[1]))).run())
