"""
Runpod 4K Upscaler Watchdog & Auto-Sync
Monitors progress on RTX 4090, auto-downloads completed 4K videos,
and terminates the pod immediately upon completion to guarantee zero extra cost.
"""

import os
import sys
import time
import json
import subprocess
import runpod

if sys.platform.startswith("win"):
    try:
        sys.stdout.reconfigure(encoding="utf-8")
        sys.stderr.reconfigure(encoding="utf-8")
    except Exception:
        pass

runpod.api_key = "rpa_5SFMBV8T5MB9XG5SSQ615M43WUUN72F1NOVEJNEU1je8b4"

SESSION_FILE = os.path.abspath(r"c:\Projects\Historical\runpod_session.json")
LOCAL_OUTPUT_DIR = os.path.abspath(r"c:\Projects\Historical\content\videos_4k")
STATE_FILE = os.path.abspath(r"c:\Projects\Historical\upscale_state.json")
SSH_KEY = os.path.expanduser(r"~/.ssh/id_ed25519")

os.makedirs(LOCAL_OUTPUT_DIR, exist_ok=True)

# Default target 6 episodes
TARGET_EPISODES = ["ep3", "ep4", "ep5", "ep6", "ep7", "ep8"]

def load_session():
    if os.path.exists(SESSION_FILE):
        with open(SESSION_FILE, "r", encoding="utf-8") as f:
            return json.load(f)
    return {}

session_info = load_session()
POD_ID = session_info.get("pod_id", "")
HOST = session_info.get("host", "")
PORT = str(session_info.get("port", "22"))
TARGET_EPISODES = session_info.get("episodes", TARGET_EPISODES)

def run_ssh(cmd_str):
    ssh_cmd = [
        "ssh", "-o", "StrictHostKeyChecking=no", "-o", "UserKnownHostsFile=/dev/null",
        "-i", SSH_KEY, "-p", PORT, f"root@{HOST}", cmd_str
    ]
    res = subprocess.run(ssh_cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, encoding="utf-8", errors="replace")
    return res.stdout.strip()

def download_file(remote_path, local_path):
    scp_cmd = [
        "scp", "-o", "StrictHostKeyChecking=no", "-P", PORT, "-i", SSH_KEY,
        f"root@{HOST}:{remote_path}", local_path
    ]
    print(f"[*] Đang tải về máy local: {os.path.basename(local_path)}...", flush=True)
    t0 = time.time()
    subprocess.run(scp_cmd, check=True)
    size_mb = os.path.getsize(local_path) / (1024 * 1024)
    print(f"✓ Đã tải xong: {os.path.basename(local_path)} ({size_mb:.1f} MB trong {time.time()-t0:.1f}s)", flush=True)

def update_local_state(completed_list):
    try:
        state = {"completed_episodes": ["ep1", "ep2"], "in_progress": None}
        if os.path.exists(STATE_FILE):
            with open(STATE_FILE, "r", encoding="utf-8") as f:
                state = json.load(f)
        all_comp = set(state.get("completed_episodes", ["ep1", "ep2"]))
        all_comp.update(completed_list)
        state["completed_episodes"] = sorted(list(all_comp))
        state["last_updated"] = time.strftime("%Y-%m-%d %H:%M:%S")
        with open(STATE_FILE, "w", encoding="utf-8") as f:
            json.dump(state, f, indent=2)
    except Exception as e:
        print(f"[-] Lỗi cập nhật state: {e}")

def kill_pod(reason=""):
    print("\n" + "=" * 60)
    print(f"🛑 [SAFETY GUARD] TIẾN HÀNH HỦY POD ĐỂ DỪNG CƯỚC NGAY LẬP TỨC: {POD_ID}")
    if reason:
        print(f"   Lý do: {reason}")
    print("=" * 60)
    terminated = False
    try:
        res = runpod.terminate_pod(POD_ID)
        print(f"✓ ĐÃ GỌI RUNPOD SDK TERMINATE: {res}")
        terminated = True
    except Exception as e:
        print(f"[-] Lỗi terminate pod qua SDK: {e}")

    # Fallback via direct GraphQL HTTP
    try:
        import urllib.request
        gql = json.dumps({"query": f'mutation {{ podTerminate(input: {{podId: "{POD_ID}"}}) }}'}).encode('utf-8')
        req = urllib.request.Request(
            f"https://api.runpod.io/graphql?api_key={runpod.api_key}",
            data=gql,
            headers={"Content-Type": "application/json"}
        )
        with urllib.request.urlopen(req, timeout=10) as r:
            resp = json.loads(r.read().decode('utf-8'))
            print(f"✓ RUNPOD GRAPHQL TERMINATE RESPONSE: {resp}")
            terminated = True
    except Exception as e:
        print(f"[-] GraphQL fallback error: {e}")

    with open(r"c:\Projects\Historical\POD_TERMINATED.txt", "w", encoding="utf-8") as f:
        f.write(f"Pod {POD_ID} terminated at {time.strftime('%Y-%m-%d %H:%M:%S')} - Reason: {reason}\n")
    print(f"✓ HOÀN TẤT THỦ TỤC TERMINATE POD {POD_ID}. CƯỚC DỪNG 100%!")

def main():
    if not POD_ID or not HOST:
        print("[-] Lỗi: Không tìm thấy thông tin session Pod trong runpod_session.json!")
        sys.exit(1)

    print("==========================================================")
    print("RUNPOD 4K PURESCALE TENSORRT WATCHDOG & COST GUARD")
    print(f"Pod ID: {POD_ID} (RTX 4090 @ Runpod)")
    print(f"Host: {HOST}:{PORT}")
    print(f"Thư mục lưu trữ: {LOCAL_OUTPUT_DIR}")
    print(f"Mục tiêu 6 tập: {TARGET_EPISODES}")
    print("==========================================================")

    start_monitor_time = time.time()
    downloaded = set()

    # Kiểm tra xem có tập nào đã tải về từ trước không
    for ep in TARGET_EPISODES:
        local_f = os.path.join(LOCAL_OUTPUT_DIR, f"{ep}_4k.mp4")
        if os.path.exists(local_f) and os.path.getsize(local_f) > 50 * 1024 * 1024:
            downloaded.add(ep)
            print(f"[*] Tập {ep} đã có sẵn tại máy local.")

    while True:
        try:
            # 1. Đọc progress.json từ Pod
            raw_json = run_ssh("cat /workspace/progress.json 2>/dev/null")
            if raw_json and raw_json.startswith("{"):
                prog = json.loads(raw_json)
                curr_ep = prog.get("current_episode")
                c_frame = prog.get("current_frame", 0)
                t_frames = prog.get("total_frames", 0)
                pct = prog.get("percent", 0)
                fps = prog.get("speed_fps", 0)
                completed_on_pod = prog.get("completed", [])

                cost_so_far = ((time.time() - start_monitor_time) / 3600) * 0.74
                if curr_ep:
                    eta_mins = (t_frames - c_frame) / max(0.1, fps) / 60
                    print(f"[{time.strftime('%H:%M:%S')}] Đang render: {curr_ep.upper()} | Frame: {c_frame}/{t_frames} ({pct}%) | {fps:.1f} FPS | ETA: {eta_mins:.1f}m | Cước dự kiến: ~${cost_so_far:.2f}", flush=True)
                else:
                    print(f"[{time.strftime('%H:%M:%S')}] Tiến độ Pod: Hoàn thành {len(completed_on_pod)}/{len(TARGET_EPISODES)} tập | Cước: ~${cost_so_far:.2f}", flush=True)

                # 2. Tải về máy local các tập đã hoàn thành trên Pod
                for ep in completed_on_pod:
                    if ep not in downloaded:
                        remote_path = f"/workspace/outputs_4k/{ep}_4k.mp4"
                        local_path = os.path.join(LOCAL_OUTPUT_DIR, f"{ep}_4k.mp4")
                        download_file(remote_path, local_path)
                        downloaded.add(ep)
                        update_local_state(list(downloaded))

                # 3. Kiểm tra điều kiện hoàn tất toàn bộ 6 tập
                if len(downloaded) >= len(TARGET_EPISODES):
                    print("\n" + "=" * 60)
                    print("🎉 TOÀN BỘ 6 TẬP ĐÃ HOÀN TẤT VÀ TẢI VỀ MÁY AN TOÀN!")
                    print("=" * 60)
                    kill_pod("Đã hoàn thành toàn bộ 6 tập.")
                    update_local_state(list(downloaded))
                    break

            # 4. Kiểm tra nếu tiến trình trt_pipeline bị dừng
            ps_check = run_ssh("ps aux | grep trt_pipeline.py | grep -v grep")
            if not ps_check:
                done_check = run_ssh("cat /workspace/ALL_DONE 2>/dev/null")
                if done_check:
                    print("✓ Tiến trình trên Pod đã báo ALL_DONE hoàn tất!")
                    # Quét tải nốt những file còn sót
                    for ep in TARGET_EPISODES:
                        if ep not in downloaded:
                            remote_path = f"/workspace/outputs_4k/{ep}_4k.mp4"
                            local_path = os.path.join(LOCAL_OUTPUT_DIR, f"{ep}_4k.mp4")
                            try:
                                download_file(remote_path, local_path)
                                downloaded.add(ep)
                            except Exception as e:
                                print(f"[-] Không thể tải {ep}: {e}")
                    kill_pod("Hoàn tất và tải xong tất cả các tập.")
                    break
                else:
                    print("[-] CẢNH BÁO: Tiến trình trt_pipeline.py bị tắt bất thường! Đang lấy log...")
                    last_log = run_ssh("tail -n 20 /workspace/trt_batch.log 2>/dev/null")
                    print(last_log)
                    kill_pod("Tiến trình bị dừng bất thường, hủy pod để bảo toàn chi phí.")
                    break

        except Exception as e:
            print(f"[-] Lỗi vòng lặp giám sát: {e}")

        time.sleep(15)

if __name__ == "__main__":
    main()
