#!/usr/bin/env python3
"""
Test script for BCT engine integration
Based on rpc_engine_test.ipynb structure
"""

import os
import time
import argparse
import multiprocessing as mp
from tqdm import tqdm

# Set CUDA device
os.environ["CUDA_VISIBLE_DEVICES"] = "4"

print("Setting up BCT engine test...")

# Import required modules
from suno_utils.gpt.generation import GenerationConfig
from suno_utils.gpt.generation_engine import make_request
from suno_utils.gpt.rpc_zmq_bct import start_bct_service_processes

# Import tracing modules
from suno_utils.tracing import start_daemon


def test_bct_engine(num_requests=4, max_gen_duration_s=30):
    # Start trace daemon first
    print("Starting trace daemon...")
    daemon_proc = mp.Process(
        target=start_daemon,
        args=(18862, 20.0, 20.0, "s3", "bct_engine_test"),
        #     port  delay duration backend project
    )
    daemon_proc.start()

    # Wait a moment for daemon to start
    time.sleep(2)

    print("Starting BCT service processes...")

    # Start BCT engine service
    client = start_bct_service_processes(
        "/app2/suno/checkpoints/2025-08-24_08-31-23/last_ckpt_infer.pt",
        max_sequences=40,
        tokenizer_path="s3://suno-data/georg/models/tokenizers/tokenizer_60k.json",
        compile=False,
        port=18865,
    )

    # Get model config
    cfg = client.get_model_cfg()
    print(f"Model config loaded: {cfg}")

    # Test lyrics (simple, non-copyrighted content)
    lyrics = """
[verse]
walking down the street today
sunshine brightening up my way
feeling good and feeling free
this is how life's meant to be

[chorus]
every step brings something new
every day a different view
life is good when you believe
in the dreams you can achieve
"""

    # Create generation config
    gconf = GenerationConfig(
        text=lyrics,
        text_tags="",
        cfg_coef=2.0,
        cfg_coef_tags=3.0,
        cfg_coef_max_steps=100,
        cfg_coef_tags_max_steps=150,
        n_batch=1,
        min_text_offset=0,
        eos_pad_duration_s=0,
        max_gen_duration_s=max_gen_duration_s,
    )

    # Create requests
    print(f"Creating {num_requests} requests...")
    import uuid

    unique_prefix = str(uuid.uuid4())[:8]
    requests = [
        make_request(f"test_{unique_prefix}_{i}", gconf, cfg, client.get_tokenizer())
        for i in range(num_requests)
    ]

    # Debug: check allow_eos setting and min generation duration
    print(f"First request allow_eos: {requests[0].allow_eos}")
    print(f"First request min_gen_duration_s: {requests[0].min_gen_duration_s}")
    print(f"First request max_gen_duration_s: {requests[0].max_gen_duration_s}")
    print(
        f"GenerationConfig min_gen_duration_s: {getattr(gconf, 'min_gen_duration_s', 'not set')}"
    )
    print(f"Semantic pad token: {cfg.semantic_pad_token}")

    # Submit jobs
    jobs = []
    print("Submitting jobs...")
    for request in tqdm(requests, desc="Adding requests"):
        job_id = client.add_request(request)
        jobs.append(job_id)
        print(f"Added job {job_id}")

    # Wait for completion
    print("Waiting for job completion...")
    start_time = time.time()

    while True:
        all_completed = True
        for job in jobs:
            state = client.get_job_state(job)
            if state is None:
                print(f"Warning: Job {job} state is None")
                continue
            if not state.get("completed", False):
                all_completed = False
                break

        if all_completed:
            print("All jobs completed!")
            break

        # Timeout after 5 minutes
        if time.time() - start_time > 300:
            print("Timeout waiting for jobs to complete")
            break

        time.sleep(1)

    # Check results
    print("\nJob results:")
    for i, job in enumerate(jobs):
        state = client.get_job_state(job)
        if state:
            completed = state["completed"]
            eos_step = state.get("eos_step")
            ttl = state.get("ttl", 0)

            print(
                f"Job {i}: completed={completed}, eos_step={eos_step}, ttl={ttl:.2f}ms"
            )

            # Get generated codes
            codes = client.get_generated_codes(job)
            if codes is not None and hasattr(codes, "shape"):
                shape = codes.shape
                print(f"  Generated codes shape: {shape}")
                # Handle both (T, 2) and (2, T) formats
                num_tokens = shape[1] if shape[0] == 2 else shape[0]
                if num_tokens > 100:
                    print(f"  ✓ Generated sufficient tokens ({num_tokens} tokens)")
                else:
                    print(
                        f"  ⚠ Generated fewer tokens than expected ({num_tokens} tokens)"
                    )
            else:
                print("  ⚠ No codes generated or invalid format")
        else:
            print(f"Job {i}: No state available")

    # Cleanup
    print("\nCleaning up...")
    try:
        client.terminate_processes()
        print("✓ BCT engine processes terminated successfully")
    except Exception as e:
        print(f"⚠ Error during cleanup: {e}")

    print("BCT engine test completed!")

    # Wait for trace daemon to finish and export
    print("Waiting for trace daemon to export traces...")
    daemon_proc.join(timeout=130)  # Wait up to 130 seconds for daemon to finish
    if daemon_proc.is_alive():
        print("Trace daemon still running, terminating...")
        daemon_proc.terminate()
        daemon_proc.join()

    print("✓ Trace export completed! Check /tmp for trace files.")
    exit(0)


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="Test BCT engine integration")
    parser.add_argument(
        "--num-requests",
        type=int,
        default=4,
        help="Number of requests to submit (default: 4)",
    )
    parser.add_argument(
        "--max-gen-duration",
        type=float,
        default=30.0,
        help="Maximum generation duration in seconds (default: 30.0)",
    )

    args = parser.parse_args()

    try:
        test_bct_engine(
            num_requests=args.num_requests, max_gen_duration_s=args.max_gen_duration
        )
    except KeyboardInterrupt:
        print("\nTest interrupted by user")
    except Exception as e:
        print(f"Test failed with error: {e}")
        import traceback

        traceback.print_exc()
