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)