#!/usr/bin/env python3
"""
Purges everything the load test created: LOADTEST-* terminals and their
status logs/items/push logs, the seeded config_file_versions rows, and the
staged OTA artifact files/directories the pushes produced. Leaves everyone
else's data untouched (everything is scoped by serial_prefix / vendor).

Usage:
    ./cleanup.py            # deletes DB rows + staged files
    ./cleanup.py --dry-run  # just prints counts, deletes nothing
"""
import argparse
import shutil
import sys
from pathlib import Path

import pymysql

from config import SETTINGS


def counts(cur):
    cur.execute(
        "SELECT COUNT(*) FROM tms_terminals WHERE serial_number LIKE %s",
        (f"{SETTINGS.serial_prefix}%",),
    )
    (terminal_count,) = cur.fetchone()

    cur.execute(
        """
        SELECT COUNT(*) FROM tms_terminal_status_logs l
        JOIN tms_terminals t ON t.id = l.terminal_id
        WHERE t.serial_number LIKE %s
        """,
        (f"{SETTINGS.serial_prefix}%",),
    )
    (log_count,) = cur.fetchone()

    cur.execute(
        "SELECT COUNT(*) FROM parameter_push_logs WHERE device_vendor = %s",
        (SETTINGS.vendor,),
    )
    (push_log_count,) = cur.fetchone()

    cur.execute(
        "SELECT id FROM config_file_versions WHERE vendor = %s",
        (SETTINGS.vendor,),
    )
    config_version_ids = [r[0] for r in cur.fetchall()]

    return terminal_count, log_count, push_log_count, config_version_ids


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--dry-run", action="store_true")
    args = parser.parse_args()

    conn = pymysql.connect(
        host=SETTINGS.db_host,
        user=SETTINGS.db_user,
        password=SETTINGS.db_password_or_exit(),
        database=SETTINGS.db_name,
        autocommit=True,
    )
    try:
        with conn.cursor() as cur:
            terminal_count, log_count, push_log_count, config_version_ids = counts(cur)
            print(f"Found: {terminal_count} terminals, {log_count} status logs, "
                  f"{push_log_count} push logs, {len(config_version_ids)} config_file_versions rows "
                  f"(vendor={SETTINGS.vendor!r}, serial prefix={SETTINGS.serial_prefix!r})")

            if args.dry_run:
                print("--dry-run: nothing deleted.")
                return

            if terminal_count == 0 and push_log_count == 0 and not config_version_ids:
                print("Nothing to clean up.")
                return

            confirm = input("Type 'delete' to purge all of the above: ")
            if confirm.strip() != "delete":
                print("Aborted.")
                return

            cur.execute("DELETE FROM parameter_push_logs WHERE device_vendor = %s", (SETTINGS.vendor,))
            print(f"  deleted {cur.rowcount} parameter_push_logs")

            cur.execute(
                """
                DELETE i FROM tms_terminal_status_items i
                JOIN tms_terminal_status_logs l ON l.id = i.status_log_id
                JOIN tms_terminals t ON t.id = l.terminal_id
                WHERE t.serial_number LIKE %s
                """,
                (f"{SETTINGS.serial_prefix}%",),
            )
            print(f"  deleted {cur.rowcount} tms_terminal_status_items")

            cur.execute(
                """
                DELETE l FROM tms_terminal_status_logs l
                JOIN tms_terminals t ON t.id = l.terminal_id
                WHERE t.serial_number LIKE %s
                """,
                (f"{SETTINGS.serial_prefix}%",),
            )
            print(f"  deleted {cur.rowcount} tms_terminal_status_logs")

            cur.execute("DELETE FROM tms_terminals WHERE serial_number LIKE %s", (f"{SETTINGS.serial_prefix}%",))
            print(f"  deleted {cur.rowcount} tms_terminals")

            if config_version_ids:
                cur.execute("DELETE FROM config_file_versions WHERE vendor = %s", (SETTINGS.vendor,))
                print(f"  deleted {cur.rowcount} config_file_versions")

        # Staged OTA artifacts: per-terminal dirs (priv/ota/LOADTEST-*), plus
        # shared-mode dirs keyed by the config_file_versions ids just deleted.
        repo = Path(SETTINGS.tms_repo_root)
        ota_dir = repo / "priv" / "ota"
        removed_dirs = 0
        for p in ota_dir.glob(f"{SETTINGS.serial_prefix}*"):
            shutil.rmtree(p, ignore_errors=True)
            removed_dirs += 1
        for cfv_id in config_version_ids:
            for shared_dir in (ota_dir / "l3" / str(cfv_id), ota_dir / "apps" / str(cfv_id)):
                if shared_dir.exists():
                    shutil.rmtree(shared_dir, ignore_errors=True)
                    removed_dirs += 1
        print(f"  removed {removed_dirs} staged OTA artifact directories under {ota_dir}")

        print("Cleanup complete.")
    finally:
        conn.close()


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