//! Complete generation flow using Tower services
//!
//! Demonstrates proper usage of sunocore Tower services:
//! 1. SubmitGeneration service
//! 2. PollGeneration service (uses GET /api/feed/v2?ids=)
//! 3. Download audio files
//!
//! Run with: cargo run --example full_generation_tower

use std::fs;
use std::path::PathBuf;
use std::time::Duration;

use bytes::Bytes;
use dotenvy::dotenv;
use futures_util::future::BoxFuture;
use http::{Request, Response, StatusCode};
use http_body_util::Full;
use sunocore::models::generation::{GenerateParams, ClipStatus};
use sunocore::services::generation::{
    SubmitGeneration, SubmitGenerationRequest,
    BatchPoll, BatchPollRequest,
};
use sunocore::validation::env::{suno_api_key_unwrap, suno_base_url_unwrap};
use tower::{Service, ServiceExt};

// Tower service adapter around reqwest::Client
#[derive(Clone)]
struct ReqwestSvc {
    client: reqwest::Client,
}

impl Service<Request<Full<Bytes>>> for ReqwestSvc {
    type Response = Response<Full<Bytes>>;
    type Error = reqwest::Error;
    type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;

    fn poll_ready(
        &mut self,
        _cx: &mut std::task::Context<'_>,
    ) -> std::task::Poll<Result<(), Self::Error>> {
        std::task::Poll::Ready(Ok(()))
    }

    fn call(&mut self, req: Request<Full<Bytes>>) -> Self::Future {
        let client = self.client.clone();
        let method = req.method().clone();
        let uri = req.uri().to_string();
        let headers = req.headers().clone();
        let body = req.into_body();

        Box::pin(async move {
            use http_body_util::BodyExt;

            let body_bytes = body.collect().await.unwrap().to_bytes();

            let mut builder = client.request(method, uri);
            for (name, value) in headers.iter() {
                builder = builder.header(name, value);
            }

            if !body_bytes.is_empty() {
                builder = builder.body(body_bytes.to_vec());
            }

            let res = builder.send().await?;
            let status = res.status();
            let body = res.bytes().await?;

            let mut resp = Response::new(Full::new(Bytes::from(body)));
            *resp.status_mut() = StatusCode::from_u16(status.as_u16()).unwrap();
            Ok(resp)
        })
    }
}

#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
    let _ = dotenv();

    let base_url = suno_base_url_unwrap();
    let api_key = suno_api_key_unwrap().ok_or("SUNO_API_KEY required")?;

    println!("=== Tower Services Generation Flow ===\n");

    // Create HTTP client and wrap in Tower service
    let client = reqwest::Client::builder()
        .timeout(Duration::from_secs(60))
        .build()?;

    let http_svc = ReqwestSvc { client: client.clone() };

    // Step 1: Submit generation using SubmitGeneration Tower service
    println!("Step 1: Submit Generation (SubmitGeneration service)");
    println!("─────────────────────────────────────────────────────");

    let mut submit_svc = SubmitGeneration::new(http_svc.clone());

    let params = GenerateParams::new("a peaceful piano melody")
        .with_title("Piano Peace")
        .with_instrumental(true);

    let submit_req = SubmitGenerationRequest::new(
        base_url.clone(),
        params,
        api_key.clone(),
    );

    let gen_response = submit_svc.ready().await?.call(submit_req).await?;

    println!("✓ Generation submitted");
    println!("  Request ID: {}", gen_response.id);
    println!("  Clips: {}", gen_response.clips.len());

    let clip_ids: Vec<String> = gen_response.clips.iter().map(|c| c.id.clone()).collect();
    for (i, id) in clip_ids.iter().enumerate() {
        println!("  Clip {}: {}", i + 1, id);
    }

    // Step 2: Poll using BatchPoll Tower service
    println!("\nStep 2: Poll for Completion (BatchPoll service)");
    println!("────────────────────────────────────────────────");
    println!("Using GET /api/feed/v2?ids=<clip-ids>\n");

    let mut poll_svc = BatchPoll::new(http_svc);

    let mut backoff_ms = 1000u64;
    let max_attempts = 30;

    let mut all_complete = false;

    for attempt in 1..=max_attempts {
        let poll_req = BatchPollRequest::new(
            base_url.clone(),
            clip_ids.clone(),
            api_key.clone(),
        );

        let statuses = poll_svc.ready().await?.call(poll_req).await?;

        print!("  Attempt {}: ", attempt);

        let mut complete_count = 0;
        for status in &statuses {
            if status.status == ClipStatus::Complete && status.audio_url.is_some() {
                print!("✓ ");
                complete_count += 1;
            } else {
                print!("○ ");
            }
        }

        println!("({}/{} ready)", complete_count, clip_ids.len());

        if complete_count >= clip_ids.len() {
            println!("\n✓ All clips complete!\n");
            all_complete = true;

            // Step 3: Download audio files
            println!("Step 3: Download Audio Files");
            println!("─────────────────────────────");

            let output_dir = PathBuf::from("/Users/neil/suno/oxide/target/generated_songs");
            fs::create_dir_all(&output_dir)?;

            let download_client = reqwest::Client::builder()
                .timeout(Duration::from_secs(60))
                .build()?;

            for (idx, status) in statuses.iter().enumerate() {
                if let Some(ref audio_url) = status.audio_url {
                    let filename = format!("{:02}_Piano_Peace.mp3", idx + 1);
                    let output_path = output_dir.join(filename);

                    print!("Downloading clip {}... ", idx + 1);

                    let resp = download_client.get(audio_url).send().await?;
                    let bytes = resp.bytes().await?;
                    fs::write(&output_path, &bytes)?;

                    println!("✓ {} bytes", bytes.len());
                }
            }

            break;
        }

        if attempt < max_attempts {
            println!("    Waiting {}ms...", backoff_ms);
            tokio::time::sleep(Duration::from_millis(backoff_ms)).await;
            backoff_ms = ((backoff_ms as f64 * 1.5) as u64).min(30000);
        }
    }

    if !all_complete {
        println!("\n! Clips not ready within {} attempts", max_attempts);
        println!("! Generation may take longer - clips are processing");
    } else {
        println!("\n━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
        println!("✓ Complete - All Tower services working!");
        println!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
    }

    Ok(())
}
