#!/usr/bin/env python3
"""Test large file transfer between two PeerDrop daemons.

Creates a 2GB test file, starts two daemons, sends the file,
and monitors memory usage throughout.
"""

import hashlib
import os
import subprocess
import sys
import time
import signal
import threading

# --- Config ---
TEST_DIR = "/tmp/peerdrop_large_test"
TEST_FILE = os.path.join(TEST_DIR, "test_2gb.bin")
RECEIVED_FILE = os.path.join(TEST_DIR, "downloads", "test_2gb.bin")
SOCK_SENDER = os.path.join(TEST_DIR, "sender.sock")
SOCK_RECEIVER = os.path.join(TEST_DIR, "receiver.sock")
PID_SENDER = os.path.join(TEST_DIR, "sender.pid")
PID_RECEIVER = os.path.join(TEST_DIR, "receiver.pid")
LOG_SENDER = os.path.join(TEST_DIR, "sender.log")
LOG_RECEIVER = os.path.join(TEST_DIR, "receiver.log")
FILE_SIZE = 2 * 1024 * 1024 * 1024  # 2GB

os.makedirs(os.path.join(TEST_DIR, "downloads"), exist_ok=True)


def get_memory_usage_mb(pid: int) -> float:
    """Get RSS memory usage of a process in MB."""
    try:
        result = subprocess.run(
            ["ps", "-o", "rss=", "-p", str(pid)],
            capture_output=True, text=True, timeout=5
        )
        rss_kb = int(result.stdout.strip())
        return rss_kb / 1024
    except Exception:
        return -1


def kill_daemon(sock_path: str, pid_path: str):
    """Kill a daemon by reading its PID file."""
    try:
        if os.path.exists(pid_path):
            with open(pid_path) as f:
                pid = int(f.read().strip())
            os.kill(pid, signal.SIGTERM)
            time.sleep(1)
            try:
                os.kill(pid, signal.SIGKILL)
            except ProcessLookupError:
                pass
            print(f"  Killed daemon PID {pid}")
        if os.path.exists(sock_path):
            os.unlink(sock_path)
        if os.path.exists(pid_path):
            os.unlink(pid_path)
    except Exception as e:
        print(f"  Cleanup error: {e}")


def wait_for_socket(sock_path: str, timeout: int = 30) -> bool:
    """Wait for a Unix socket file to appear."""
    for i in range(timeout):
        time.sleep(1)
        if os.path.exists(sock_path):
            return True
    return False


def main():
    print("=" * 60)
    print("PeerDrop Large File Transfer Test (2GB)")
    print("=" * 60)

    # Step 1: Create 2GB test file
    print(f"\n[1/6] Creating {FILE_SIZE // (1024*1024)}MB test file...")
    if os.path.exists(TEST_FILE) and os.path.getsize(TEST_FILE) == FILE_SIZE:
        print(f"  Test file already exists, skipping creation")
    else:
        chunk = os.urandom(1024 * 1024)  # 1MB random data
        bytes_written = 0
        with open(TEST_FILE, "wb") as f:
            while bytes_written < FILE_SIZE:
                write_size = min(len(chunk), FILE_SIZE - bytes_written)
                f.write(chunk[:write_size])
                bytes_written += write_size
                if bytes_written % (200 * 1024 * 1024) == 0:
                    print(f"  Written {bytes_written // (1024*1024)}MB / {FILE_SIZE // (1024*1024)}MB")
        print(f"  Created: {os.path.getsize(TEST_FILE)} bytes")

    # Compute hash of original file
    print("  Computing SHA256 of original file...")
    sha256_orig = hashlib.sha256()
    with open(TEST_FILE, "rb") as f:
        while True:
            data = f.read(8 * 1024 * 1024)  # 8MB chunks
            if not data:
                break
            sha256_orig.update(data)
    orig_hash = sha256_orig.hexdigest()
    print(f"  Original SHA256: {orig_hash}")

    # Step 2: Kill any existing daemons
    print("\n[2/6] Cleaning up any existing daemons...")
    kill_daemon(SOCK_SENDER, PID_SENDER)
    kill_daemon(SOCK_RECEIVER, PID_RECEIVER)

    # Step 3: Start receiver daemon
    print("\n[3/6] Starting receiver daemon (port 4002)...")
    receiver_proc = subprocess.Popen(
        [sys.executable, "-m", "peerdrop.interfaces.cli.app",
         "--sock", SOCK_RECEIVER,
         "daemon", "start",
         "--port", "4002",
         "--download-dir", os.path.join(TEST_DIR, "downloads")],
        stdout=open(LOG_RECEIVER, "w"),
        stderr=subprocess.STDOUT,
    )
    print(f"  Receiver PID: {receiver_proc.pid}")

    if not wait_for_socket(SOCK_RECEIVER):
        print("  ERROR: Receiver daemon did not start in 30s")
        return 1
    print("  Receiver ready")

    # Step 4: Start sender daemon
    print("\n[4/6] Starting sender daemon (port 4001)...")
    sender_proc = subprocess.Popen(
        [sys.executable, "-m", "peerdrop.interfaces.cli.app",
         "--sock", SOCK_SENDER,
         "daemon", "start",
         "--port", "4001"],
        stdout=open(LOG_SENDER, "w"),
        stderr=subprocess.STDOUT,
    )
    print(f"  Sender PID: {sender_proc.pid}")

    if not wait_for_socket(SOCK_SENDER):
        print("  ERROR: Sender daemon did not start in 30s")
        return 1
    print("  Sender ready")

    # Step 5: Get identities and connect
    print("\n[5/6] Connecting sender to receiver...")

    # Use the IPC client directly for more reliable communication
    sys.path.insert(0, os.getcwd())
    from peerdrop.interfaces.client import PeerDropClient

    sender_client = PeerDropClient(SOCK_SENDER)
    sender_client.connect()
    sender_resp = sender_client.get_identity()
    sender_data = sender_resp.get("data", {})
    sender_peer_id = sender_data.get("peer_id", "")
    sender_addrs = sender_data.get("addrs", [])
    print(f"  Sender Peer ID: {sender_peer_id}")
    print(f"  Sender Addrs: {sender_addrs[:2]}...")

    receiver_client = PeerDropClient(SOCK_RECEIVER)
    receiver_client.connect()
    receiver_resp = receiver_client.get_identity()
    receiver_data = receiver_resp.get("data", {})
    receiver_peer_id = receiver_data.get("peer_id", "")
    receiver_addrs = receiver_data.get("addrs", [])
    print(f"  Receiver Peer ID: {receiver_peer_id}")
    print(f"  Receiver Addrs: {receiver_addrs[:2]}...")

    # Connect receiver to sender using first IPv4 address
    connected = False
    for addr in sender_addrs:
        if "/ip4/" in addr and "/ip6/" not in addr:
            print(f"  Connecting receiver to sender at {addr}...")
            try:
                receiver_client.connect_peer(addr)
                connected = True
                print("  Connected!")
                break
            except Exception as e:
                print(f"  Connection failed: {e}")

    if not connected:
        print("  WARNING: Could not connect via IPv4, trying all addresses...")
        for addr in sender_addrs[:1]:
            try:
                receiver_client.connect_peer(addr)
                connected = True
                print("  Connected!")
                break
            except Exception as e:
                print(f"  Connection failed: {e}")

    time.sleep(2)  # Let connection settle

    # Step 6: Send the file
    print(f"\n[6/6] Sending 2GB file from sender to receiver...")
    print(f"  File: {TEST_FILE} ({FILE_SIZE // (1024*1024)}MB)")
    print(f"  Target: {receiver_peer_id}")

    # Monitor memory in background
    memory_log = []
    monitoring = True

    def monitor_memory():
        while monitoring:
            sender_mem = get_memory_usage_mb(sender_proc.pid)
            receiver_mem = get_memory_usage_mb(receiver_proc.pid)
            memory_log.append({
                "time": time.time(),
                "sender_mb": sender_mem,
                "receiver_mb": receiver_mem,
            })
            if len(memory_log) % 5 == 0:
                print(f"  [Memory] Sender: {sender_mem:.1f}MB, Receiver: {receiver_mem:.1f}MB")
            time.sleep(1)

    monitor_thread = threading.Thread(target=monitor_memory, daemon=True)
    monitor_thread.start()

    start_time = time.time()

    try:
        transfer = sender_client.send_file(TEST_FILE, receiver_peer_id)
        print(f"  Transfer result: {transfer}")
    except Exception as e:
        print(f"  Transfer error: {e}")
        import traceback
        traceback.print_exc()

    elapsed = time.time() - start_time
    monitoring = False
    monitor_thread.join(timeout=2)

    # Step 7: Verify received file
    print("\n" + "=" * 60)
    print("RESULTS")
    print("=" * 60)

    if os.path.exists(RECEIVED_FILE):
        received_size = os.path.getsize(RECEIVED_FILE)
        print(f"\n  Received file: {RECEIVED_FILE}")
        print(f"  Expected size: {FILE_SIZE}")
        print(f"  Actual size:   {received_size}")
        print(f"  Size match: {'YES' if received_size == FILE_SIZE else 'NO'}")

        if received_size == FILE_SIZE:
            print("  Computing SHA256 of received file...")
            sha256_recv = hashlib.sha256()
            with open(RECEIVED_FILE, "rb") as f:
                while True:
                    data = f.read(8 * 1024 * 1024)
                    if not data:
                        break
                    sha256_recv.update(data)
            recv_hash = sha256_recv.hexdigest()
            print(f"  Received SHA256: {recv_hash}")
            print(f"  Hash match: {'YES' if recv_hash == orig_hash else 'NO'}")
    else:
        print(f"\n  ERROR: Received file not found at {RECEIVED_FILE}")
        # Check download dir
        download_dir = os.path.join(TEST_DIR, "downloads")
        if os.path.exists(download_dir):
            files = os.listdir(download_dir)
            print(f"  Files in download dir: {files}")

    print(f"\n  Transfer time: {elapsed:.1f}s")
    speed = FILE_SIZE / elapsed / (1024 * 1024) if elapsed > 0 else 0
    print(f"  Average speed: {speed:.1f} MB/s")

    # Memory stats
    if memory_log:
        max_sender = max(m["sender_mb"] for m in memory_log)
        max_receiver = max(m["receiver_mb"] for m in memory_log)
        file_mb = FILE_SIZE / (1024*1024)
        print(f"\n  Peak sender memory:   {max_sender:.1f}MB")
        print(f"  Peak receiver memory: {max_receiver:.1f}MB")
        print(f"  File size:            {file_mb:.1f}MB")
        print(f"  Sender memory/file ratio: {max_sender / file_mb:.2f}x")
        print(f"  Receiver memory/file ratio: {max_receiver / file_mb:.2f}x")

        if max_sender < file_mb * 2 and max_receiver < file_mb * 2:
            print("\n  PASS: Memory usage is within acceptable bounds (< 2x file size)")
        else:
            print("\n  FAIL: Memory usage too high (> 2x file size)")

    # Cleanup
    print("\nCleaning up...")
    sender_client.close()
    receiver_client.close()
    kill_daemon(SOCK_SENDER, PID_SENDER)
    kill_daemon(SOCK_RECEIVER, PID_RECEIVER)

    return 0


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