#!/usr/bin/python3
# @file
# @brief CII OLDB KVDB snapshot tool
#
# @copyright
#   SPDX-FileCopyrightText: 2026 European Southern Observatory (ESO) @n
#   SPDX-License-Identifier: LGPL-3.0-only

import subprocess
import sys
import os
from textwrap import dedent
import time
import argparse
from concurrent.futures import ThreadPoolExecutor, as_completed
import shutil

# ----------------------------
# Documentation

THIS = os.path.basename(sys.argv[0])

what = "Manage saved content of an OLDB KVDB instance (single or cluster)"

usage = f"""Usage: {THIS} [-p port] save|unsave, or -h"""

help = f"""
Commands:
  save  : create snapshot (will be auto-reloaded on KVDB start)
  unsave: remove snapshot to prevent auto-reloading - IMPORTANT

This tool allows to manually manage saved state of a KVDB instance
for use in cases where the KVDB has automatic state saving disabled.

Notes:
  1. The KVDB supports only 1 snapshot (saved state) at any time,
     i.e. existing saved state gets replaced whenever you "save".
  2. For using this utility, the KVDB must be currently running.

Important:
  If automatic state saving is disabled on the KVDB, and you
  manually create a snapshot using this tool, you are responsible
  for removing it after use. As long as you don't remove it, the
  KVDB will, on EVERY future start, be reset to your snapshot.
"""

def print_workflow(title, done, fail):
    c = ["0", "1", "2", "3", "4", "5", "6"]
    for i in done: c[i] = "✓"
    for i in fail: c[i] = "✗"
    print(dedent(f"""
    {title}
      [{c[0]}] KVDB is running
      [{c[1]}] $ {THIS} save
      [{c[2]}] Stop the KVDB  (via cii-services, systemd, etc.)
      [{c[3]}] Perform the maintenance activities you want to do
      [{c[4]}] Start the KVDB (via cii-services, systemd, etc.)
      [{c[5]}] KVDB loads the saved state
      [{c[6]}] $ {THIS} unsave
    """)[1:-1])

# ----------------------------
# Defaults

MAX_WORKERS = 4
SEED_PORT = 6379
CLI_CLIENT = shutil.which("valkey-cli") or shutil.which("redis-cli") or None

# ----------------------------
# Helpers

def run_cli(host, port, *args, check=False):
    if CLI_CLIENT is None: raise RuntimeError("A KVDB client (like valkey-cli) is required, but none was found")
    cmd = [CLI_CLIENT, "-h", host, "-p", str(port), *args]
    return subprocess.run(cmd, capture_output=True, text=True, check=check)

def wait_for_bgsave(host, port, timeout=120):
    start = time.time()

    while True:
        res = run_cli(host, port, "INFO", "persistence")

        if "rdb_bgsave_in_progress:0" in res.stdout:
            return True

        if time.time() - start > timeout:
            raise RuntimeError(f"BGSAVE timeout on {host}:{port}")

        time.sleep(1)

def echo (*str):
    print (f"{THIS}:", *str)

# ---------------------------------------
# Dynamic discovery of kdvb setup

def is_accessible(host, port):
    res = run_cli(host, port, "INFO")
    return not "not connect" in res.stderr

def is_cluster(host, port):
    res = run_cli(host, port, "INFO", "cluster")
    return "cluster_enabled:1" in res.stdout

def get_node_props(host, port):
    res = run_cli(host, port, "CONFIG", "GET", "dir", "cluster-config-file", "dbfilename", "save")
    out = res.stdout.strip().splitlines()
    ori = dict(zip(out[0::2], out[1::2]))
    return {
        'host': host,
        'port': port,
        'dir' : ori['dir'],
        'conf': ori['cluster-config-file'],
        'dump': ori['dbfilename'],
        'save': ori['save'],
    }

def discover_nodes(seed_host="127.0.0.1", seed_port=6379):

    if not is_accessible(seed_host, seed_port):
        return None

    if not is_cluster(seed_host, seed_port):
        host,port = seed_host,seed_port
        return get_node_props (host, port)

    # cluster

    res = run_cli (seed_host, seed_port, "CLUSTER", "NODES", check=True)
    nodes = []
    known = set()

    for line in res.stdout.strip().splitlines():
        parts = line.split()

        addr = parts[1]   # host:port@bus-port
        flags = parts[2]

        # skip disconnected nodes
        if "fail" in flags:
            continue
        
        # pending: for increased performance, consider:
        # - deploying a replica for each master
        # - backing up the replicas only (not the masters)
        # note:
        #   we are not using replicas currently,
        #   so i'm leaving this here just as an example
        if "master" not in flags:
            continue

        host_port = addr.split("@")[0]
        host,port = host_port.split(":")

        # de-duplication of redundant results
        if host_port in known:
            continue
        known.add(host_port)

        node = get_node_props (host, port)
        nodes.append(node)

    return nodes


# ----------------------------
# Internal

def save_node (node):
    host = node["host"]
    port = node["port"]

    echo(f"[{port}] Triggering BGSAVE")
    run_cli(host, port, "BGSAVE")

    echo(f"[{port}] Waiting for BGSAVE")
    wait_for_bgsave(host, port)

    echo(f"[{port}] Created saved state")
    wait_for_bgsave(host, port)

    return port

def get_file_path(node, filetype, check_exist=True):
    ret = None
    if filetype == "conf": ret = node["conf"]
    if filetype == "dump": ret = os.path.join(node["dir"], node["dump"])
    if check_exist and ret and not os.path.exists(ret):
        ret = None
    return ret

def unsave_node(node):

    # remove  dump file
    if f:= get_file_path(node, "dump"):
        os.remove(f)

    # remove cluster config file
    if f:= get_file_path(node, "conf"):
        os.remove(f)

    port = node["port"]
    echo(f"[{port}] Removed saved state")

# ----------------------------
# User Commands

def save():

    nodes = discover_nodes(seed_port=SEED_PORT)
    if not nodes:
        print_workflow("Problem occurred:", done=[], fail=[0])
        return

    todos = [node["port"] for node in nodes]

    with ThreadPoolExecutor(max_workers=MAX_WORKERS) as executor:
        futures = [executor.submit(save_node, node) for node in nodes]

        for future in as_completed(futures):
            try:
                port = future.result()
                todos.remove(port)
            except Exception as e:
                echo(f"Failure while saving state: {e}")

    if todos:
        for port in todos:
            echo(f"{port} State could not be saved")
        print_workflow("Problem occurred:", done=[0], fail=[1])
        return

    # Some guiding for occasional users

    print("")
    print_workflow("You are here:", done=[0,1], fail=[])

    auto_saves = [s for n in nodes if (s := n["save"])]
    # for auto-save to work, all nodes must have it enabled
    if len(auto_saves) == len(nodes):
        print (dedent(f"""
        Note: your KVDB is configured for automatic state saving, so
        there should be no need for you to manage saved state yourself.
        State was saved as you requested, but feel free to skip "unsave":
          [6] <unsave not strictly needed>"""))

    dump_files = [f for n in nodes if (f := get_file_path(n,"dump"))]
    conf_files = [f for n in nodes if (f := get_file_path(n,"conf"))]
    print (dedent(f"""
    Note: in case you change your mind, and want to get rid of the saved
    state while the KVDB is stopped (the unsave command requires the KVDB
    to be running), run this command before starting the KVDB:
      [3] $ rm -f {" ".join(dump_files)} {" ".join(conf_files)}
      [4] Start the KVDB (via cii-services, systemd, etc.)
      [5] <nothing gets reloaded, no saved state>
      [6] <unsave not needed, no saved state>
    """))


def unsave():

    nodes = discover_nodes(seed_port=SEED_PORT)
    if not nodes:
        print_workflow("Problem occurred:", done=[], fail=[0])
        return

    for node in nodes:
        unsave_node(node)


# ----------------------------
# Main

def main():

    if "-h" in sys.argv or "--help" in sys.argv:
        print(what + "\n" + usage + "\n" + help)
        print_workflow("Workflow:", [], [])
        exit(0)

    parser = argparse.ArgumentParser(description=what, add_help=False)
    parser.add_argument("-p", type=int, help="seed port", dest="port")
    subparsers = parser.add_subparsers(dest="command")
    subparsers.add_parser("save", help="Create snapshot")
    subparsers.add_parser("unsave", help="Remove snapshot")
    args = parser.parse_args()

    if args.port is not None:
        global SEED_PORT
        SEED_PORT = args.port

    match args.command:
        case "save": save()
        case "unsave": unsave()
        case None: parser.print_help() # no command → show help / fallback
        case _: echo("Invalid command:", args.command)

if __name__ == "__main__":
    main()
