#!/usr/bin/env python3
"""Run a four-node OpenThread simulation and retain the actual CLI transcript."""
import argparse
import json
import os
import pty
import re
import select
import subprocess
import sys
import time
from pathlib import Path

NETWORK_KEY = "00112233445566778899aabbccddeeff"  # local teaching network only
BINARY_VERSION = "OPENTHREAD/thread-reference-20250612"


class Node:
    def __init__(self, name, node_id, binary):
        self.name = name
        master, slave = pty.openpty()
        self.master = master
        self.proc = subprocess.Popen([str(binary), str(node_id)], stdin=slave,
                                     stdout=slave, stderr=slave, close_fds=True)
        os.close(slave)
        self.read_until_prompt(5)

    def read_until_prompt(self, timeout):
        data = b""
        deadline = time.monotonic() + timeout
        while time.monotonic() < deadline:
            ready, _, _ = select.select([self.master], [], [], 0.1)
            if ready:
                try:
                    data += os.read(self.master, 65536)
                except OSError as exc:
                    raise RuntimeError(f"{self.name} exited ({self.proc.poll()}): {exc}") from exc
                if data.endswith(b"> "):
                    return data.decode(errors="replace").replace("\r", "")
        raise TimeoutError(f"{self.name}: CLI did not finish within {timeout}s: {data[-250:]!r}")

    def cmd(self, command, timeout=10):
        os.write(self.master, (command + "\n").encode())
        output = self.read_until_prompt(timeout)
        if "Error " in output or "InvalidCommand" in output:
            raise RuntimeError(f"{self.name}: {command}: {output.strip()}")
        return output.strip().removesuffix("> ").rstrip()

    def close(self):
        if self.proc.poll() is None:
            self.proc.terminate()
            try:
                self.proc.wait(timeout=3)
            except subprocess.TimeoutExpired:
                self.proc.kill()
                self.proc.wait()
        os.close(self.master)


def answer(output):
    lines = [line for line in output.splitlines() if line not in ("Done", "> ")]
    return lines[1].strip() if len(lines) > 1 else ""


def wait_state(node, expected, timeout=60):
    deadline = time.monotonic() + timeout
    while time.monotonic() < deadline:
        value = answer(node.cmd("state"))
        if value == expected:
            return value
        time.sleep(1)
    raise TimeoutError(f"{node.name} did not become {expected}; last state={value}")


def wait_parent_change(node, former, timeout=80):
    deadline = time.monotonic() + timeout
    last = ""
    while time.monotonic() < deadline:
        state = answer(node.cmd("state"))
        if state == "child":
            output = node.cmd("parent")
            last = output
            match = re.search(r"^Ext Addr: ([0-9a-f]+)$", output, re.M)
            if match and match.group(1) != former:
                return output
        else:
            last = state
        time.sleep(2)
    raise TimeoutError(f"child did not select a new parent after relay loss: {last}")


def run(binary, output_dir, capture_steps=False):
    output_dir.mkdir(parents=True, exist_ok=True)
    nodes = {}
    transcript = []
    snapshots = []
    shown_step = None
    capture_commands = {
        1: {("leader", "dataset active"), ("leader", "state"), ("leader", "ipaddr mleid")},
        2: {("relay-a", "state"), ("leader", "router table")},
        3: {("child", "state"), ("child", "parent"), ("relay-a", "child table")},
        4: {("relay-b", "state"), ("child", "parent"), ("child", "ping")},
        5: {("relay-a", "thread stop"), ("relay-a", "state"), ("child", "parent"),
            ("child", "state"), ("relay-b", "child table")},
        6: {("child", "ping"), ("child", "state"), ("leader", "router table")},
    }

    def pause_after_step():
        if capture_steps:
            sys.stdout.flush()
            sys.stdin.readline()

    def command(name, value, section=None):
        nonlocal shown_step
        result = nodes[name].cmd(value, timeout=20)
        block = f"[{name}] $ {value}\n{result}"
        transcript.append(block)
        (output_dir / "cli-transcript.txt").write_text("\n\n".join(transcript) + "\n")
        if section:
            snapshots.append({"step": section, "block": block})
        if capture_steps and section and any(
            name == node and (value == wanted or (wanted == "ping" and value.startswith("ping ")))
            for node, wanted in capture_commands[section]
        ):
            if shown_step != section:
                print("\033[2J\033[H", end="")
                print(f"OpenThread simulated CLI  |  step {section}  |  live ot-cli-ftd run\n")
                shown_step = section
            lines = result.splitlines()
            if lines and lines[0] == value:
                lines.pop(0)
            lines = [line for line in lines if line not in ("Done", ">", "> ")]
            if section == 1 and value == "dataset active":
                lines = [line for line in lines if line.startswith((
                    "Active Timestamp:", "Channel:", "Ext PAN ID:", "Mesh Local Prefix:",
                    "Network Name:", "PAN ID:"))]
            print(f"[{name}] $ {value}", *lines, sep="\n", flush=True)
        return result

    try:
        for name, node_id in (("leader", 1), ("relay-a", 2), ("relay-b", 3), ("child", 4)):
            nodes[name] = Node(name, node_id, binary)
        version = answer(command("leader", "version", 1))
        if not version.startswith(BINARY_VERSION):
            raise RuntimeError(f"expected {BINARY_VERSION}, got {version}")
        ext = {name: answer(command(name, "extaddr")) for name in nodes}
        edges = {
            "leader": ["relay-a", "relay-b"],
            "relay-a": ["leader", "child"],
            "relay-b": ["leader", "child"],
            "child": ["relay-a", "relay-b"],
        }
        for name, neighbors in edges.items():
            for other in neighbors:
                command(name, f"macfilter addr add {ext[other]}")
            command(name, "macfilter addr allowlist")
        command("child", "routereligible disable")
        command("child", "childtimeout 20")
        for name in ("leader", "relay-a", "relay-b"):
            command(name, "routerselectionjitter 1")

        command("leader", "dataset init new", 1)
        command("leader", "dataset channel 15", 1)
        command("leader", "dataset panid 0x1234", 1)
        command("leader", "dataset extpanid 1122334455667788", 1)
        command("leader", "dataset networkname ThreadRouteLab", 1)
        command("leader", f"dataset networkkey {NETWORK_KEY}", 1)
        command("leader", "dataset meshlocalprefix fd12:3456:789a:1::", 1)
        command("leader", "dataset commit active", 1)
        command("leader", "dataset active", 1)
        dataset = answer(command("leader", "dataset active -x", 1))
        for name in ("relay-a", "relay-b", "child"):
            command(name, f"dataset set active {dataset}")
        command("leader", "ifconfig up", 1)
        command("leader", "thread start", 1)
        wait_state(nodes["leader"], "leader")
        command("leader", "state", 1)
        command("leader", "ipaddr mleid", 1)
        pause_after_step()

        command("relay-a", "ifconfig up", 2)
        command("relay-a", "thread start", 2)
        wait_state(nodes["relay-a"], "router")
        command("relay-a", "state", 2)
        command("leader", "router table", 2)
        pause_after_step()

        command("child", "ifconfig up", 3)
        command("child", "thread start", 3)
        wait_state(nodes["child"], "child")
        command("child", "state", 3)
        before_parent = command("child", "parent", 3)
        if ext["relay-a"] not in before_parent:
            raise RuntimeError("baseline child did not attach to relay-a")
        command("relay-a", "child table", 3)
        pause_after_step()

        command("relay-b", "ifconfig up", 4)
        command("relay-b", "thread start", 4)
        wait_state(nodes["relay-b"], "router")
        command("relay-b", "state", 4)
        command("leader", "router table", 4)
        command("child", "parent", 4)
        leader_address = answer(command("leader", "ipaddr mleid", 4))
        first_ping = command("child", f"ping {leader_address} 16 1", 4)
        if "1 packets received" not in first_ping:
            raise RuntimeError(f"baseline ICMPv6 probe failed: {first_ping}")
        time.sleep(2)
        # Ping replies can arrive after Done; capture the next prompt's queued bytes.
        command("child", "state", 4)
        pause_after_step()

        command("relay-a", "thread stop", 5)
        command("relay-a", "state", 5)
        changed_parent = wait_parent_change(nodes["child"], ext["relay-a"])
        if ext["relay-b"] not in changed_parent:
            raise RuntimeError("child selected an unexpected parent")
        command("child", "parent", 5)
        command("child", "state", 5)
        new_children = command("relay-b", "child table", 5)
        if ext["child"] not in new_children:
            raise RuntimeError("alternate relay does not list the child")
        command("leader", "router table", 5)
        pause_after_step()
        second_ping = command("child", f"ping {leader_address} 16 1", 6)
        if "1 packets received" not in second_ping:
            raise RuntimeError(f"post-fault ICMPv6 probe failed: {second_ping}")
        time.sleep(2)
        command("child", "state", 6)
        command("leader", "router table", 6)
        pause_after_step()

        (output_dir / "cli-transcript.txt").write_text("\n\n".join(transcript) + "\n")
        (output_dir / "snapshots.json").write_text(json.dumps(snapshots, indent=2) + "\n")
        print(f"PASS: {version}")
        print(f"PASS: leader={ext['leader']} relay-a={ext['relay-a']} relay-b={ext['relay-b']} child={ext['child']}")
        print("PASS: child attached to relay-a, then relay-b after relay-a stopped")
        print(f"PASS: transcript={output_dir / 'cli-transcript.txt'}")
    finally:
        for node in nodes.values():
            node.close()


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--binary", type=Path,
                        default=Path("openthread-thread-reference-20250612/build/simulation/examples/apps/cli/ot-cli-ftd"),
                        help="pinned ot-cli-ftd executable (defaults to the README build path)")
    parser.add_argument("--output-dir", type=Path, default=Path("thread-route-output"))
    parser.add_argument("--capture-steps", action="store_true",
                        help="show live step output in a terminal and wait for Enter after each step")
    args = parser.parse_args()
    if not args.binary.is_file():
        parser.error(f"OpenThread CLI binary not found at {args.binary}; follow README.md to build it")
    run(args.binary.resolve(), args.output_dir, args.capture_steps)
