from fastapi import FastAPI, Request, Depends, HTTPException, status, BackgroundTasks
from fastapi.middleware.cors import CORSMiddleware
from fastapi.templating import Jinja2Templates
from fastapi.responses import HTMLResponse, StreamingResponse, JSONResponse, Response
from fastapi.staticfiles import StaticFiles
import os
import logging
import json
import boto3
from pathlib import Path
from dotenv import load_dotenv
import httpx
import asyncio
from app.models import *
from app.suno_service import SunoService
from app.process_lyrics import process_lyrics_simple
from pydantic import BaseModel
from typing import Optional, AsyncGenerator
from fastapi.security import HTTPBasic, HTTPBasicCredentials
import secrets
from datetime import datetime, timezone
import re
import threading
from PIL import Image, ImageDraw, ImageFont
import io


# Load environment variables
load_dotenv()

# Configure logging
logging.basicConfig(
    level=os.getenv("LOG_LEVEL", "INFO"),
    format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
    handlers=[logging.StreamHandler(), logging.FileHandler("app.log")],
)
logger = logging.getLogger(__name__)


# Update the templates directory path to be relative to the current file
BASE_DIR = Path(__file__).resolve().parent
templates = Jinja2Templates(directory=str(BASE_DIR / "templates"))

app = FastAPI()

# Add CORS middleware
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)

# Mount static files directory
app.mount("/static", StaticFiles(directory="data"), name="static")

# Add new S3 client for lyrics
lyrics_s3_client = boto3.client(
    "s3",
    aws_access_key_id=os.getenv("AWS_ACCESS_KEY_ID"),
    aws_secret_access_key=os.getenv("AWS_SECRET_ACCESS_KEY"),
    region_name=os.getenv("AWS_REGION", "us-east-2"),
)

# Add security scheme
security = HTTPBasic()


# Add authentication function
def authenticate(credentials: HTTPBasicCredentials = Depends(security)):
    correct_password = "AlexasTheBest!"
    is_correct_password = secrets.compare_digest(
        credentials.password.encode("utf8"), correct_password.encode("utf8")
    )

    if not is_correct_password:
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="Incorrect password",
            headers={"WWW-Authenticate": "Basic"},
        )
    return credentials


@app.get("/", response_class=HTMLResponse)
async def root(
    request: Request, credentials: HTTPBasicCredentials = Depends(authenticate)
):
    return templates.TemplateResponse("index.html", {"request": request})


@app.get("/create_song_page", response_class=HTMLResponse)
async def create_song_page(
    request: Request, credentials: HTTPBasicCredentials = Depends(authenticate)
):
    """
    Render the song creation page.
    """
    return templates.TemplateResponse("create_song_minimal.html", {"request": request})


@app.get("/oauth_test", response_class=HTMLResponse)
async def oauth_test_page(
    request: Request, credentials: HTTPBasicCredentials = Depends(authenticate)
):
    """
    Render the OAuth test page.
    """
    return templates.TemplateResponse("oauth_test.html", {"request": request})


@app.get("/oauth_callback", response_class=HTMLResponse)
async def oauth_callback_page(
    request: Request,
    code: str = None,
    state: str = None,
    error: str = None,
    error_description: str = None,
    credentials: HTTPBasicCredentials = Depends(authenticate),
):
    """
    Render the OAuth callback page.
    This page receives and displays the authorization code and other parameters from Suno.
    """
    return templates.TemplateResponse(
        "oauth_callback.html",
        {
            "request": request,
            "code": code,
            "state": state,
            "error": error,
            "error_description": error_description,
        },
    )


# Helper to parse SSE lines
def parse_sse_event(lines: list[str]) -> tuple[Optional[str], Optional[str]]:
    event_name = None
    event_data = ""
    for line in lines:
        if line.startswith("event:"):
            event_name = line[len("event:") :].strip()
        elif line.startswith("data:"):
            event_data += line[len("data:") :].strip()
        # Ignore empty lines and comments
    # Use 'message' as default event name if not specified
    return event_name if event_name else "message", event_data if event_data else None


# Helper to listen to a Suno SSE stream and put events on a queue
async def listen_suno_sse(
    url: str,
    source_name: str,
    queue: asyncio.Queue,
    client: httpx.AsyncClient,
    relevant_events: set[str],
):
    buffer = []
    reconnection_delay = 1  # Initial delay in seconds
    max_reconnection_delay = 60

    while True:  # Add retry logic
        try:
            async with client.stream("GET", url, timeout=None) as response:
                response.raise_for_status()  # Raise exception for 4xx/5xx status
                reconnection_delay = 1  # Reset delay on successful connection
                logger.info(f"Connected to Suno {source_name} stream: {url}")
                await queue.put(
                    (
                        "backend_status",
                        json.dumps(
                            {"source": source_name, "message": f"Connected to {url}"}
                        ),
                    )
                )

                async for line in response.aiter_lines():
                    if not line:  # Empty line signifies end of an event
                        if buffer:
                            event_name, event_data = parse_sse_event(buffer)
                            buffer = []  # Reset buffer for next event
                            if event_name in relevant_events and event_data:
                                await queue.put((event_name, event_data))
                            elif (
                                event_name and event_data
                            ):  # Log other events if needed for debugging
                                logger.debug(
                                    f"Ignoring Suno event: {source_name} ({event_name})"
                                )
                    else:
                        buffer.append(line)

                # Stream finished normally
                logger.info(f"Suno {source_name} stream finished.")
                await queue.put(
                    (
                        "backend_status",
                        json.dumps(
                            {"source": source_name, "message": "Stream finished"}
                        ),
                    )
                )
                return  # Exit retry loop on clean finish

        except httpx.RequestError as e:
            logger.warning(
                f"Connection error for {source_name} stream {url}: {e}. Retrying in {reconnection_delay}s..."
            )
            await queue.put(
                (
                    "backend_error",
                    json.dumps(
                        {
                            "source": source_name,
                            "error": f"Connection error: {e}. Retrying...",
                        }
                    ),
                )
            )
        except httpx.HTTPStatusError as e:
            logger.error(
                f"HTTP error for {source_name} stream {url}: {e.response.status_code} - {e.response.text}. Retrying in {reconnection_delay}s..."
            )
            await queue.put(
                (
                    "backend_error",
                    json.dumps(
                        {
                            "source": source_name,
                            "error": f"HTTP error {e.response.status_code}. Retrying...",
                        }
                    ),
                )
            )
        except Exception as e:
            logger.error(
                f"Unexpected error in {source_name} listener for {url}: {e}. Retrying in {reconnection_delay}s...",
                exc_info=True,
            )
            await queue.put(
                (
                    "backend_error",
                    json.dumps(
                        {
                            "source": source_name,
                            "error": f"Unexpected listener error. Retrying...",
                        }
                    ),
                )
            )

        await asyncio.sleep(reconnection_delay)
        reconnection_delay = min(
            reconnection_delay * 2, max_reconnection_delay
        )  # Exponential backoff


# Define request models
class CreateSongRequest(BaseModel):
    topic: str
    tags: Optional[str] = None
    model: Optional[str] = None


# Define latency log file and lock
LATENCY_LOG_FILE = "latency_log.jsonl"
latency_log_lock = threading.Lock()


# Helper function to parse ISO timestamp safely
def parse_iso_timestamp(timestamp_str: Optional[str]) -> Optional[datetime]:
    if not timestamp_str:
        return None
    try:
        dt = datetime.fromisoformat(timestamp_str.replace("Z", "+00:00"))
        if dt.tzinfo is None:
            dt = dt.replace(tzinfo=timezone.utc)
        return dt
    except ValueError:
        logger.warning(f"Could not parse timestamp: {timestamp_str}")
        return None


# Helper function to calculate latency in seconds
def calculate_latency_seconds(
    end_dt: Optional[datetime], start_dt: Optional[datetime]
) -> Optional[float]:
    if end_dt and start_dt:
        latency = end_dt - start_dt
        return latency.total_seconds()
    return None


@app.get("/latency_data")
async def get_latency_data(credentials: HTTPBasicCredentials = Depends(authenticate)):
    """
    Reads latency log data, calculates specific latency metrics,
    and returns the raw values for frontend processing.
    """
    latencies = {
        "time_to_title": [],
        "time_to_lyrics": [],
        "time_to_image": [],
        "time_to_audio": [],
        "image_gen_latency": [],
        "audio_gen_latency": [],
    }

    try:
        with latency_log_lock:  # Use lock for thread-safe reading
            if not os.path.exists(LATENCY_LOG_FILE):
                logger.warning(f"{LATENCY_LOG_FILE} not found.")
                return JSONResponse(
                    content={
                        "error": f"{LATENCY_LOG_FILE} not found.",
                        "data": latencies,
                    },
                    status_code=404,
                )

            with open(LATENCY_LOG_FILE, "r") as f:
                for line in f:
                    try:
                        log_entry = json.loads(line)
                        milestones = log_entry.get("milestones", {})
                        model = log_entry.get("model", "unknown")

                        if model != "chirp-v4-h-api":
                            logger.debug(
                                f"Skipping entry for clip {log_entry.get('clip_id')}: model ({model}) != chirp-v4-h-api"
                            )
                            continue

                        t0 = parse_iso_timestamp(milestones.get("t0_start"))
                        t2 = parse_iso_timestamp(milestones.get("t2_title"))
                        t3 = parse_iso_timestamp(milestones.get("t3_lyrics"))
                        t4 = parse_iso_timestamp(milestones.get("t4_image"))
                        t5 = parse_iso_timestamp(milestones.get("t5_audio_stream"))

                        if not t0:  # Skip entry if t0 is missing/invalid
                            continue

                        # Calculate latencies only if timestamps are valid
                        t2_t0 = calculate_latency_seconds(t2, t0)

                        # Skip this entry if time to title is > 10 seconds
                        if t2_t0 is not None and t2_t0 > 10:
                            logger.debug(
                                f"Skipping entry for clip {log_entry.get('clip_id')}: time_to_title ({t2_t0:.2f}s) > 10s"
                            )
                            continue

                        t3_t0 = calculate_latency_seconds(t3, t0)
                        t4_t0 = calculate_latency_seconds(t4, t0)
                        t5_t0 = calculate_latency_seconds(t5, t0)
                        t4_t3 = (
                            calculate_latency_seconds(t4, t3) if t3 else None
                        )  # Need t3 for these
                        t5_t3 = (
                            calculate_latency_seconds(t5, t3) if t3 else None
                        )  # Need t3 for these

                        # Skip this entry if time to audio is > 20 seconds
                        if t5_t0 is not None and t5_t0 > 20:
                            logger.debug(
                                f"Skipping entry for clip {log_entry.get('clip_id')}: time_to_audio ({t5_t0:.2f}s) > 20s"
                            )
                            continue

                        # Skip this entry if time to image is > 20 seconds
                        if t4_t0 is not None and t4_t0 > 20:
                            logger.debug(
                                f"Skipping entry for clip {log_entry.get('clip_id')}: time_to_image ({t4_t0:.2f}s) > 20s"
                            )
                            continue

                        # Append valid latencies to lists
                        if t2_t0 is not None:
                            latencies["time_to_title"].append(t2_t0)
                        if t3_t0 is not None:
                            latencies["time_to_lyrics"].append(t3_t0)
                        if t4_t0 is not None:
                            latencies["time_to_image"].append(t4_t0)
                        if t5_t0 is not None:
                            latencies["time_to_audio"].append(t5_t0)
                        if t4_t3 is not None:
                            latencies["image_gen_latency"].append(t4_t3)
                        if t5_t3 is not None:
                            latencies["audio_gen_latency"].append(t5_t3)

                    except json.JSONDecodeError as e:
                        logger.warning(
                            f"Skipping invalid JSON line in {LATENCY_LOG_FILE}: {e}"
                        )
                    except Exception as e:
                        logger.error(
                            f"Unexpected error processing line in {LATENCY_LOG_FILE}: {e}",
                            exc_info=True,
                        )

        return JSONResponse(content={"data": latencies})

    except Exception as e:
        logger.error(f"Error reading or processing latency log: {e}", exc_info=True)
        return JSONResponse(
            content={
                "error": f"Failed to process latency data: {str(e)}",
                "data": latencies,
            },
            status_code=500,
        )


@app.post("/create_song")
async def create_song(
    request: CreateSongRequest,
    background_tasks: BackgroundTasks,
    credentials: HTTPBasicCredentials = Depends(authenticate),
):
    """
    Create a song using Suno's API, connect to Suno's SSE streams server-side,
    and forward relevant events to the client over a single SSE stream.
    """

    async def event_generator() -> AsyncGenerator[str, None]:
        start_time = datetime.utcnow()  # Use UTC time
        server_milestones = {
            "t0_start": start_time.isoformat(),  # Record start time immediately
            "t2_title": None,
            "t3_lyrics": None,
            "t4_image": None,
            "t5_audio_stream": None,
        }
        clip_id = None
        request_id = None
        suno_service = None
        tasks = []
        queue = asyncio.Queue()
        listener_tasks_count = 2  # We expect two listener tasks (clip, request)
        finished_listeners = 0
        line_count_for_title = 0  # Specific counter for title detection
        final_status = "unknown"  # Track final status for logging
        error_details = None  # Store error details for logging
        # Flags for early exit condition
        received_image_generated = False
        received_gen_streaming = False

        # Helper to format SSE data payload with timestamp and conditional milestones
        def format_sse_data(event_name: str, data: dict) -> str:
            now_iso = datetime.utcnow().isoformat()
            payload = {
                "server_timestamp": now_iso,
                # "server_milestones": server_milestones, # Include current milestones conditionally
                "original_data": data,  # Nest original data
            }
            # Only include milestones for non-'line' events
            if event_name != "line":
                payload["server_milestones"] = server_milestones
            return json.dumps(payload)

        # Helper to yield a complete SSE event string
        def format_sse_event(event_name: str, data: dict) -> str:
            json_data = format_sse_data(event_name, data)
            return f"event: {event_name}\ndata: {json_data}\n\n"

        try:
            # 1. Initial Setup & Send Start Event
            api_key = os.getenv("SUNO_API_KEY")
            if not api_key:
                error_data = {"error": "SUNO_API_KEY not configured"}
                yield format_sse_event("error", error_data)
                return

            suno_service = SunoService(api_key=api_key)
            start_data = {"message": "Song generation requested"}
            yield format_sse_event("start", start_data)
            logger.info("Create song request received.")

            # 2. Initiate Song Generation
            generation_response = await suno_service.generate_song(
                topic=request.topic, tags=request.tags, model=request.model
            )
            # Log the raw response for debugging
            logger.debug(f"Raw Suno generation response: {generation_response}")

            # Note: creation_time is now derived from milestones server-side
            clip_id = str(generation_response.id)
            request_id = getattr(generation_response, "request_id", None)

            if not clip_id or not request_id:
                error_msg = "Failed to get clip_id or request_id from Suno."
                logger.error(error_msg + f" Response: {generation_response}")
                error_data = {"error": error_msg, "detail": str(generation_response)}
                yield format_sse_event("error", error_data)
                return

            logger.info(
                f"Song generation initiated: clip_id={clip_id}, request_id={request_id}"
            )

            # 3. Send Created Event
            created_data = {
                "clip_id": clip_id,
                "request_id": request_id,
                "status": generation_response.status,
            }
            yield format_sse_event("created", created_data)

            # 4. Connect to Suno SSE Streams Concurrently
            SUNO_SSE_BASE = "https://audiopipe.suno.ai"
            clip_event_url = f"{SUNO_SSE_BASE}/clip_events/?clip_id={clip_id}"
            request_event_url = (
                f"{SUNO_SSE_BASE}/request_events/?request_id={request_id}"
            )

            # Events we care about from each stream
            clip_relevant_events = {
                "image_generated",
                "gen_streaming",
                "metadata_update",
            }
            request_relevant_events = {"line", "lyrics", "generate_queued"}

            async with httpx.AsyncClient() as client:
                # Create tasks for listening to each stream
                task1 = asyncio.create_task(
                    listen_suno_sse(
                        clip_event_url, "clip", queue, client, clip_relevant_events
                    )
                )
                task2 = asyncio.create_task(
                    listen_suno_sse(
                        request_event_url,
                        "request",
                        queue,
                        client,
                        request_relevant_events,
                    )
                )
                tasks = [task1, task2]

                # 5. Consume events from queue and forward to client
                try:
                    while finished_listeners < listener_tasks_count:
                        event_name, event_data_json = await queue.get()
                        queue.task_done()  # Mark task as done immediately

                        original_event_data = {}
                        is_line_event = event_name == "line"
                        try:
                            # Parse the original event data from Suno/backend
                            if is_line_event:
                                # Line event data is a JSON string, parse it to get raw string
                                parsed_line = json.loads(event_data_json)
                                # Convert None to empty string, otherwise use parsed value
                                original_event_data = (
                                    parsed_line if parsed_line is not None else ""
                                )
                            elif event_data_json and (
                                event_data_json.startswith("{")
                                or event_data_json.startswith("[")
                            ):
                                # Other events are expected to be JSON objects/arrays
                                original_event_data = json.loads(event_data_json)
                            else:
                                # Handle unexpected non-JSON, non-line data
                                original_event_data = {"raw_data": event_data_json}

                        except json.JSONDecodeError:
                            logger.warning(
                                f"Failed to parse original JSON data for event {event_name}: {event_data_json}"
                            )
                            # For lines, send the raw (quoted) string on parse failure?
                            if is_line_event:  # Check if it's a line event
                                original_event_data = (
                                    event_data_json  # Send raw quoted string
                                )
                            else:
                                original_event_data = {
                                    "error": "Failed to parse original data",
                                    "raw_data": event_data_json,
                                }

                        # --- Milestone Tracking & Event Forwarding ---
                        now_utc_iso = datetime.utcnow().isoformat()

                        if event_name == "backend_status":
                            # Log internal status but don't forward raw status to client by default
                            source = original_event_data.get("source", "unknown")
                            message = original_event_data.get("message", "")
                            logger.info(f"Backend Status ({source}): {message}")
                            if message == "Stream finished":
                                finished_listeners += 1
                            # Optionally forward a sanitized status event if needed later
                            # yield format_sse_event("backend_info", {"source": source, "message": message})

                        elif event_name == "backend_error":
                            # Log internal error and forward a sanitized error event to client
                            source = original_event_data.get("source", "unknown")
                            error_msg = original_event_data.get(
                                "error", "Unknown backend error"
                            )
                            logger.error(f"Backend Error ({source}): {error_msg}")
                            yield format_sse_event(
                                "error",
                                {
                                    "error": f"Backend stream error [{source}]",
                                    "detail": error_msg,
                                },
                            )
                            # Consider if we should break or continue retrying based on error type

                        # --- Suno Event Processing ---
                        elif event_name == "line":
                            line_count_for_title += 1
                            if (
                                line_count_for_title == 1
                                and server_milestones["t2_title"] is None
                            ):
                                server_milestones["t2_title"] = now_utc_iso
                            # Forward the line event with current milestones
                            yield format_sse_event(event_name, original_event_data)

                        elif event_name == "lyrics":
                            if server_milestones["t3_lyrics"] is None:
                                server_milestones["t3_lyrics"] = now_utc_iso
                            yield format_sse_event(event_name, original_event_data)

                        elif event_name == "image_generated":
                            if server_milestones["t4_image"] is None:
                                server_milestones["t4_image"] = now_utc_iso
                            yield format_sse_event(event_name, original_event_data)
                            received_image_generated = True  # Set flag

                        elif event_name == "gen_streaming":
                            if server_milestones["t5_audio_stream"] is None:
                                server_milestones["t5_audio_stream"] = now_utc_iso
                            yield format_sse_event(event_name, original_event_data)
                            received_gen_streaming = True  # Set flag

                        else:
                            # Forward other relevant Suno events
                            yield format_sse_event(event_name, original_event_data)

                        # Check for early exit condition
                        if received_image_generated and received_gen_streaming:
                            logger.info(
                                "Received both image and audio stream start events. Exiting early."
                            )
                            final_status = (
                                "aborted_early_audio_image"  # Set specific status
                            )
                            break  # Exit the loop
                except asyncio.CancelledError:
                    logger.info(
                        "Event generator task cancelled, likely due to client disconnect or server shutdown."
                    )
                    # The finally block will handle listener task cleanup
                    raise  # Re-raise the error to ensure FastAPI knows the stream ended prematurely

                logger.info("Both Suno SSE listeners have finished.")

            # 6. Final Completion Event
            end_time = datetime.utcnow()
            total_elapsed_time = (end_time - start_time).total_seconds()
            complete_data = {
                "message": "Processing complete.",
                "total_elapsed_time": total_elapsed_time,
            }
            # Ensure final milestones are included
            # Only yield complete event if we didn't exit early
            if final_status != "aborted_early_audio_image":
                yield format_sse_event("complete", complete_data)
                final_status = "complete"  # Mark as complete for logging
            logger.info("Finished forwarding events to client.")

        except httpx.HTTPStatusError as e:
            logger.error(
                f"Initial Suno API call failed: {e.response.status_code} - {e.response.text}",
                exc_info=True,
            )
            error_time = datetime.utcnow()
            elapsed_time = (error_time - start_time).total_seconds()
            error_data = {
                "error": f"Suno API error: {e.response.status_code}",
                "detail": e.response.text,
                "total_elapsed_time": elapsed_time,
            }
            final_status = "error"  # Mark as error
            error_details = error_data  # Store error details
            yield format_sse_event("error", error_data)
        except Exception as e:
            logger.error(f"Error in create_song event generator: {e}", exc_info=True)
            error_time = datetime.utcnow()
            elapsed_time = (
                (error_time - start_time).total_seconds() if start_time else None
            )
            error_data = {"error": str(e), "total_elapsed_time": elapsed_time}
            final_status = "error"  # Mark as error
            error_details = error_data  # Store error details
            yield format_sse_event("error", error_data)
        finally:
            # Ensure background tasks are cancelled if the generator exits prematurely
            for task in tasks:
                if not task.done():
                    task.cancel()
            # Wait for tasks to finish cancellation
            if tasks:
                await asyncio.gather(*tasks, return_exceptions=True)

            # Log latency data using BackgroundTasks
            log_entry = {
                "log_timestamp": datetime.utcnow().isoformat(),
                "clip_id": clip_id,
                "request_id": request_id,
                "topic": request.topic,
                "model": request.model,
                "status": final_status,
                "milestones": server_milestones,
                "error_details": error_details,  # Will be null if status is complete
            }

            # Define the logging function to run in background
            def write_log(entry):
                try:
                    with latency_log_lock:  # Acquire lock for thread-safe file writing
                        with open(LATENCY_LOG_FILE, "a") as f:
                            json.dump(entry, f)
                            f.write("\n")  # Add newline for JSON Lines format
                    logger.info(
                        f"Latency data logged for clip_id: {entry.get('clip_id')}"
                    )
                except Exception as log_err:
                    logger.error(
                        f"Failed to write latency log for clip_id {entry.get('clip_id')}: {log_err}"
                    )

            background_tasks.add_task(write_log, log_entry)  # Schedule the task

            logger.info("Event generator finished cleanup.")

    return StreamingResponse(event_generator(), media_type="text/event-stream")


@app.get("/public/lyrics/{clip_id}")
async def get_public_lyrics(
    clip_id: str,
    format: str = "plain",  # Options: plain, line, word, hoot
):
    """
    Public endpoint to get song lyrics in various formats.

    Parameters:
    - clip_id: The song's unique identifier
    - format: The desired format (plain, line, word, hoot)

    Returns lyrics in the requested format if available.
    """
    try:
        # Step 1: Check if song exists via song status
        api_key = os.getenv("SUNO_API_KEY")
        if not api_key:
            return JSONResponse(
                content={"error": "API key not configured"}, status_code=500
            )

        suno_service = SunoService(api_key=api_key)
        song_response = await suno_service.get_song_status(song_id=clip_id)

        if not song_response:
            return JSONResponse(
                content={"error": f"No song found with ID: {clip_id}"}, status_code=404
            )

        # For "plain" format, try to use the lyrics from metadata.prompt if available
        if (
            format == "plain"
            and song_response.metadata
            and song_response.metadata.prompt
        ):
            # Get lyrics from metadata.prompt and clean up section markers
            lyrics_text = song_response.metadata.prompt
            # Remove section markers like [Verse] or [Chorus]
            lyrics_text = re.sub(r"\[\w+(?:\s*\d*)?\]\s*", "", lyrics_text)

            # Create a simple WebVTT from plain text
            plain_vtt = "WEBVTT\n\n"
            plain_vtt += lyrics_text

            return Response(
                content=plain_vtt,
                media_type="text/vtt",
            )

        # If we need time-aligned lyrics, the song must be complete
        if format != "plain" and song_response.status != "complete":
            return JSONResponse(
                content={
                    "error": f"Song generation not complete: {song_response.status}"
                },
                status_code=400,
            )

        # For plain format, if we got here it means metadata.prompt wasn't available
        if format == "plain" and song_response.status != "complete":
            return JSONResponse(
                content={
                    "error": f"Song generation not complete and no lyrics available: {song_response.status}"
                },
                status_code=400,
            )

        # At this point, we know the song is complete, so hoot.json should be available
        try:
            s3_key = f"studio/uploads/{clip_id}_hoot.json"
            response = lyrics_s3_client.get_object(
                Bucket="suno-data-uploads", Key=s3_key
            )

            hoot_content = response["Body"].read().decode("utf-8")
            hoot_data = json.loads(hoot_content)
        except lyrics_s3_client.exceptions.NoSuchKey:
            # If we get here, it means the song is complete but hoot.json is missing
            # This is unexpected since complete songs should have hoot.json
            return JSONResponse(
                content={
                    "error": f"Aligned lyrics data not available for completed clip {clip_id}"
                },
                status_code=404,
            )

        # Process lyrics using the process_lyrics_simple function
        plain_text, line_timestamps_vtt, word_timestamps_vtt = process_lyrics_simple(
            hoot_data
        )

        # Return requested format
        if format == "plain":
            # Create a simple WebVTT from plain text
            plain_vtt = "WEBVTT\n\n"
            # plain_vtt += "00:00:00.000 --> 99:59:59.999\n"
            plain_vtt += plain_text

            return Response(
                content=plain_vtt,
                media_type="text/vtt",
                headers={
                    "Content-Disposition": f'attachment; filename="{clip_id}_plain.vtt"'
                },
            )
        elif format == "line":
            return Response(
                content=line_timestamps_vtt,
                media_type="text/vtt",
            )
        elif format == "word":
            return Response(
                content=word_timestamps_vtt,
                media_type="text/vtt",
            )
        elif format == "hoot":
            return JSONResponse(
                content=hoot_data,
            )
        else:
            return JSONResponse(
                content={
                    "error": f"Invalid format: {format}. Use 'plain', 'line', 'word', or 'hoot'"
                },
                status_code=400,
            )

    except Exception as e:
        logger.error(f"Error fetching lyrics for clip {clip_id}: {str(e)}")
        return JSONResponse(content={"error": str(e)}, status_code=500)


@app.get("/public/image/")
async def get_custom_album_art(tags: str = ""):
    """
    Public endpoint to get custom album art with tags as text overlay.

    Parameters:
    - tags: Comma-separated list of tags (e.g., "pop,rock,catchy")

    Returns a PNG image with tags overlaid as text.
    """
    try:
        # Load the base image
        image_path = os.path.join(BASE_DIR, "static", "album_art.png")
        img = Image.open(image_path)

        # Set up drawing context
        draw = ImageDraw.Draw(img)

        # Load font (using default if specific font not available)
        try:
            font = ImageFont.truetype("Arial", 64)
        except IOError:
            font = ImageFont.load_default().font_variant(size=64)

        # Process tags (limit to first 3)
        tag_list = tags.split(",")[:3] if tags else []

        # Draw tags on image, one per line
        y_position = 100  # Starting Y position
        for tag in tag_list:
            draw.text((100, y_position), tag.strip(), fill="white", font=font)
            y_position += 100  # Move down for next tag

        # Convert the modified image to bytes
        img_byte_arr = io.BytesIO()
        img.save(img_byte_arr, format="PNG")
        img_byte_arr.seek(0)

        # Return the image
        return Response(content=img_byte_arr.getvalue(), media_type="image/png")
    except Exception as e:
        logger.error(f"Error generating custom album art: {str(e)}")
        return JSONResponse(content={"error": str(e)}, status_code=500)
