import argparse
import libtmux


def get_node_addrs():
    with open("/etc/hosts", "r") as fi:
        res = fi.readlines()
    res = [r.strip() for r in res]
    res = [r for r in res if r]
    ii = res.index("# END ANSIBLE MANAGED BLOCK BASTION")
    assert res[ii + 1].startswith("# BEGIN ANSIBLE MANAGED BLOCK")
    hosts = res[ii + 2 :]
    hosts = [h.split()[1].replace("-rdma.local.rdma", "") for h in hosts if not h.startswith("#")]
    hosts = list(set(hosts))
    assert len(hosts) <= 8
    return hosts


def get_session(name):
    server = libtmux.Server()
    sessions = server.sessions.filter(name=name)
    assert len(sessions) == 1, sessions
    return sessions[0]


def create_and_attach_new_window(session):
    wn = session.new_window(window_name="new_window")
    return wn


def split_window(window):
    vertical = False
    attach = True
    for _ in range(3):
        for pane in window.panes:
            pane = pane.split_window(vertical=vertical, attach=attach)
        vertical = not vertical
        attach = not attach
    return window


def put_text_to_panes(window, text, enter=True):
    if isinstance(text, str):
        text = {str(i): text for i in range(len(window.panes))}
    for pane in window.panes:
        t = text[pane.index]
        pane.send_keys(t, enter=enter)


def create_window_with_ssh(session_name, new_window_name=None):
    if new_window_name is None:
        new_window_name = "new_window"
    session = get_session(session_name)
    window = create_and_attach_new_window(session)
    window = split_window(window)
    hosts = get_node_addrs()
    hosts = sorted(hosts)
    cmds = [f"ssh {host}" for host in hosts]

    text_dict = {str(i): h for i, h in enumerate(cmds)}
    put_text_to_panes(window, text_dict)
    return window


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--mode", help="create or put ", type=str, required=True)
    parser.add_argument("--session_name", help="name of tmux session", type=str, required=True)
    parser.add_argument(
        "--window_name",
        help="display a square of a given number",
        type=str,
        default=None,
    )
    parser.add_argument("--command", default=None)
    args = parser.parse_args()

    session = get_session(args.session_name)
    if args.mode == "create":
        window = create_window_with_ssh(args.session_name, new_window_name=args.window_name)
        text_dict = {str(pane.index): f"export PANE_ID={pane.index}" for pane in window.panes}
        put_text_to_panes(window, text_dict)
        if args.command is not None:
            put_text_to_panes(window, args.command)
    elif args.mode == "put":
        window = session.windows.filter(window_name=args.window_name)[0]
        put_text_to_panes(window, args.command)
    else:
        raise ValueError()
