760 lines
26 KiB
Python
760 lines
26 KiB
Python
"""
|
|
helpers.py
|
|
==========
|
|
Utility and configuration helpers for the Theme Song Booth Flask app.
|
|
|
|
These functions are stateless (or use Flask's current_app / session context)
|
|
and are imported by app.py. Keeping them here reduces the size of the route file.
|
|
"""
|
|
|
|
import os
|
|
import re
|
|
import shutil
|
|
import smtplib
|
|
import ssl
|
|
import time
|
|
import base64
|
|
import hmac
|
|
import hashlib
|
|
import secrets
|
|
import json
|
|
from email.message import EmailMessage
|
|
from pathlib import Path
|
|
|
|
from flask import session, current_app, flash
|
|
from werkzeug.utils import secure_filename
|
|
|
|
from cryptography.fernet import Fernet
|
|
from cryptography.hazmat.primitives import hashes
|
|
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
|
|
|
|
from mutagen.mp3 import MP3
|
|
from mutagen.easyid3 import EasyID3
|
|
|
|
import requests
|
|
|
|
from models import update_request
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Genre / decade helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
_GENRES_PATH = Path('/mnt/Storage/Music Genres.txt')
|
|
_DECADES_PATH = Path('/mnt/Storage/Decades.txt')
|
|
_FALLBACK_GENRES_PATH = Path(__file__).parent / 'lists' / 'music_genres.txt'
|
|
_FALLBACK_DECADES_PATH = Path(__file__).parent / 'lists' / 'decades.txt'
|
|
|
|
|
|
def _load_lines(path: Path) -> list[str]:
|
|
"""Load a text file and return non-empty stripped lines."""
|
|
if not path.exists():
|
|
return []
|
|
lines = path.read_text(encoding='utf-8').splitlines()
|
|
return [line.strip() for line in lines if line.strip()]
|
|
|
|
|
|
def _load_list(primary: Path, fallback: Path) -> list[str]:
|
|
"""Load from the primary path, falling back to the bundled copy."""
|
|
lines = _load_lines(primary)
|
|
if lines:
|
|
return lines
|
|
return _load_lines(fallback)
|
|
|
|
|
|
MUSIC_GENRES = _load_list(_GENRES_PATH, _FALLBACK_GENRES_PATH)
|
|
DECADES = _load_list(_DECADES_PATH, _FALLBACK_DECADES_PATH)
|
|
|
|
|
|
def parse_style_genre(style_genre: str | None) -> dict:
|
|
"""
|
|
Split a stored combined style string into decade, basic, and additional.
|
|
The stored format is 'Decade, Basic, Additional' (additional may be empty).
|
|
"""
|
|
parts = [p.strip() for p in (style_genre or '').split(',') if p.strip()]
|
|
return {
|
|
'decade': parts[0] if len(parts) > 0 else '',
|
|
'basic_style': parts[1] if len(parts) > 1 else '',
|
|
'additional_style': ', '.join(parts[2:]) if len(parts) > 2 else '',
|
|
}
|
|
|
|
|
|
def build_style_genre(decade: str, basic_style: str, additional_style: str) -> str:
|
|
"""Build the combined style_genre string stored in the database."""
|
|
parts = [p.strip() for p in [decade, basic_style, additional_style] if p.strip()]
|
|
return ', '.join(parts)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Auth / validation helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def is_admin():
|
|
"""Return True if the current browser session is logged in as admin."""
|
|
return session.get('admin') is True
|
|
|
|
|
|
def require_admin():
|
|
"""Redirect to the admin login page if the user is not logged in."""
|
|
from flask import redirect, url_for
|
|
if not is_admin():
|
|
return redirect(url_for('admin_login'))
|
|
|
|
|
|
def admin_password_ok(pw):
|
|
"""Check the submitted admin password against the configured one."""
|
|
return pw and pw == current_app.config['ADMIN_PASSWORD']
|
|
|
|
|
|
def is_valid_email(email):
|
|
"""Return True if the given string looks like a valid email address."""
|
|
if not email:
|
|
return False
|
|
pattern = r"^[^\s@]+@[^\s@]+\.[^\s@]+$"
|
|
return re.match(pattern, email) is not None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# File upload helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def allowed_file(filename):
|
|
"""Return True if the uploaded filename has an allowed extension (mp3)."""
|
|
return '.' in filename and filename.rsplit('.', 1)[1].lower() in current_app.config['ALLOWED_EXTENSIONS']
|
|
|
|
|
|
def upload_path(request_id):
|
|
"""Return the per-request upload directory path, creating it if necessary."""
|
|
p = Path(current_app.config['UPLOAD_FOLDER']) / str(request_id)
|
|
p.mkdir(parents=True, exist_ok=True)
|
|
return p
|
|
|
|
|
|
def save_upload(request_id, file_obj, version, song_title=None):
|
|
"""
|
|
Save an uploaded MP3 file for a request, preserving the original filename
|
|
with a version prefix. Applies the configured metadata tags.
|
|
"""
|
|
if not file_obj or file_obj.filename == '':
|
|
return None
|
|
if not allowed_file(file_obj.filename):
|
|
flash('Only MP3 files are allowed.', 'error')
|
|
return None
|
|
original = secure_filename(file_obj.filename)
|
|
filename = f"{version.upper()} - {original}"
|
|
p = upload_path(request_id)
|
|
dest = p / filename
|
|
file_obj.save(dest)
|
|
apply_mp3_tags(str(dest), song_title)
|
|
return str(dest)
|
|
|
|
|
|
def apply_mp3_tags(path, title=None):
|
|
"""
|
|
Write common ID3 tags on an uploaded MP3 using the runtime metadata defaults.
|
|
Failures are logged as a warning and do not block the upload.
|
|
"""
|
|
cfg = load_booth_settings()
|
|
try:
|
|
audio = MP3(path)
|
|
if audio.tags is None:
|
|
audio.add_tags()
|
|
if not isinstance(audio.tags, EasyID3):
|
|
audio.tags = EasyID3()
|
|
tags = audio.tags
|
|
if title:
|
|
tags['title'] = title
|
|
if cfg.get('artist'):
|
|
tags['artist'] = cfg['artist']
|
|
if cfg.get('album'):
|
|
tags['album'] = cfg['album']
|
|
if cfg.get('year'):
|
|
tags['date'] = str(cfg['year'])
|
|
audio.save()
|
|
if cfg.get('comment'):
|
|
from mutagen.id3 import COMM, TXXX
|
|
audio2 = MP3(path)
|
|
if audio2.tags is None:
|
|
audio2.add_tags()
|
|
audio2.tags["COMM"] = COMM(encoding=3, lang='eng', desc='Comment', text=cfg['comment'])
|
|
audio2.tags["TXXX:Comment"] = TXXX(encoding=3, desc='Comment', text=cfg['comment'])
|
|
audio2.save()
|
|
except Exception as e:
|
|
try:
|
|
flash(f'Warning: could not tag MP3: {e}', 'error')
|
|
except RuntimeError:
|
|
import logging
|
|
logging.getLogger('app').warning('Could not tag MP3 %s: %s', path, e)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Encryption / settings helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def _get_fernet():
|
|
"""Derive a Fernet key from the Flask SECRET_KEY so stored values are encrypted."""
|
|
secret = current_app.config['SECRET_KEY'].encode()
|
|
kdf = PBKDF2HMAC(
|
|
algorithm=hashes.SHA256(),
|
|
length=32,
|
|
salt=b'theme-song-booth-v1',
|
|
iterations=480000,
|
|
)
|
|
key = base64.urlsafe_b64encode(kdf.derive(secret))
|
|
return Fernet(key)
|
|
|
|
|
|
def encrypt_value(value):
|
|
"""Encrypt a string using the Flask SECRET_KEY. Returns base64 ciphertext."""
|
|
if not value:
|
|
return ''
|
|
return _get_fernet().encrypt(value.encode()).decode()
|
|
|
|
|
|
def decrypt_value(ciphertext):
|
|
"""Decrypt a string previously encrypted by encrypt_value."""
|
|
if not ciphertext:
|
|
return ''
|
|
try:
|
|
return _get_fernet().decrypt(ciphertext.encode()).decode()
|
|
except Exception:
|
|
return ''
|
|
|
|
|
|
def decrypt_value_legacy(ciphertext):
|
|
"""Decrypt or return plaintext. Tolerates unencrypted legacy values."""
|
|
if not ciphertext:
|
|
return ''
|
|
plaintext = decrypt_value(ciphertext)
|
|
if plaintext:
|
|
return plaintext
|
|
if not ciphertext.endswith('='):
|
|
return ciphertext
|
|
return ''
|
|
|
|
|
|
def settings_file_path():
|
|
"""Return the path to the persistent runtime settings JSON file."""
|
|
return Path(current_app.config['DATABASE']).parent / current_app.config['SETTINGS_FILE']
|
|
|
|
|
|
def load_booth_settings():
|
|
"""Load persistent runtime settings from JSON file inside the upload parent."""
|
|
cfg_path = settings_file_path()
|
|
if cfg_path.exists():
|
|
try:
|
|
cfg = json.loads(cfg_path.read_text())
|
|
for key in ('artist', 'album', 'year', 'comment', 'smtp_host', 'smtp_port', 'smtp_user', 'smtp_from'):
|
|
if cfg.get(key) is None:
|
|
cfg[key] = ''
|
|
return cfg
|
|
except (json.JSONDecodeError, OSError):
|
|
pass
|
|
return {}
|
|
|
|
|
|
def save_booth_settings(settings):
|
|
"""Persist runtime settings to JSON file."""
|
|
cfg_path = settings_file_path()
|
|
try:
|
|
cfg_path.write_text(json.dumps(settings, indent=2))
|
|
except OSError as e:
|
|
flash(f'Warning: could not save settings: {e}', 'error')
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Config getters
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def get_email_config():
|
|
"""Return the effective SMTP configuration."""
|
|
cfg = load_booth_settings()
|
|
return {
|
|
'SMTP_HOST': cfg.get('smtp_host', current_app.config['SMTP_HOST']),
|
|
'SMTP_PORT': int(cfg.get('smtp_port') or current_app.config['SMTP_PORT']),
|
|
'SMTP_USER': cfg.get('smtp_user', current_app.config['SMTP_USER']),
|
|
'SMTP_PASS': decrypt_value(cfg.get('smtp_pass', '')) or current_app.config['SMTP_PASS'],
|
|
'SMTP_FROM': cfg.get('smtp_from', current_app.config['SMTP_FROM']),
|
|
}
|
|
|
|
|
|
def get_refresh_seconds():
|
|
"""Return the dashboard auto-refresh interval in seconds (10, 20, or 30)."""
|
|
cfg = load_booth_settings()
|
|
try:
|
|
val = int(cfg.get('refresh_seconds', 10))
|
|
except (ValueError, TypeError):
|
|
val = 10
|
|
return val if val in (10, 20, 30) else 10
|
|
|
|
|
|
def get_kiosk_cycle_seconds():
|
|
"""Return the kiosk slide cycle interval in seconds."""
|
|
cfg = load_booth_settings()
|
|
try:
|
|
val = int(cfg.get('kiosk_cycle_seconds', 10))
|
|
except (ValueError, TypeError):
|
|
val = 10
|
|
if val == 0:
|
|
return 0
|
|
return max(5, val)
|
|
|
|
|
|
def get_kiosk_mode():
|
|
"""Return 'qr', 'prices', 'queue', or 'cycle' based on kiosk_cycle_seconds setting."""
|
|
cfg = load_booth_settings()
|
|
try:
|
|
val = int(cfg.get('kiosk_cycle_seconds', 10))
|
|
except (ValueError, TypeError):
|
|
val = 10
|
|
if val == -1:
|
|
return 'qr'
|
|
if val == 0:
|
|
return 'prices'
|
|
if val == 1:
|
|
return 'queue'
|
|
return 'cycle'
|
|
|
|
|
|
def get_max_revisions():
|
|
"""Return the effective max revisions as an integer."""
|
|
cfg = load_booth_settings()
|
|
try:
|
|
val = int(cfg.get('max_revisions', current_app.config.get('MAX_REVISIONS', 2)))
|
|
except (ValueError, TypeError):
|
|
val = current_app.config.get('MAX_REVISIONS', 2)
|
|
return max(0, val)
|
|
|
|
|
|
def get_callback_expiry_hours():
|
|
"""Return the Hermes signed callback token lifetime in hours (default 168 = 7 days)."""
|
|
cfg = load_booth_settings()
|
|
try:
|
|
val = int(cfg.get('callback_expiry_hours', 168))
|
|
except (ValueError, TypeError):
|
|
val = 168
|
|
return max(1, val)
|
|
|
|
|
|
def get_musicgpt_api_key():
|
|
"""Return the MusicGPT API key from env only."""
|
|
return current_app.config.get('MUSICGPT_API_KEY', '')
|
|
|
|
|
|
def get_musicgpt_default_model():
|
|
"""Return the configured default MusicGPT model."""
|
|
default = current_app.config.get('MUSICGPT_DEFAULT_MODEL', 'v6-pro')
|
|
models = get_musicgpt_models()
|
|
return default if default in models else models[-1]
|
|
|
|
|
|
def get_musicgpt_models():
|
|
"""Return the list of supported MusicGPT models."""
|
|
return list(current_app.config.get('MUSICGPT_MODELS', ['v6', 'v6-pro']))
|
|
|
|
|
|
def build_musicgpt_webhook_url():
|
|
"""Build the public webhook URL for MusicGPT async callbacks."""
|
|
base = current_app.config.get('PUBLIC_BASE_URL', '').rstrip('/')
|
|
return f"{base}/api/musicgpt/webhook"
|
|
|
|
|
|
def get_hermes_api_key():
|
|
"""Return the effective Hermes API key."""
|
|
env_key = current_app.config.get('HERMES_API_KEY', '')
|
|
if env_key:
|
|
return env_key
|
|
cfg = load_booth_settings()
|
|
return decrypt_value_legacy(cfg.get('hermes_api_key', ''))
|
|
|
|
|
|
def set_hermes_api_key(key):
|
|
"""Persist a new Hermes API key (encrypted) to runtime settings."""
|
|
cfg = load_booth_settings()
|
|
cfg['hermes_api_key'] = encrypt_value(key)
|
|
save_booth_settings(cfg)
|
|
|
|
|
|
def generate_hermes_api_key():
|
|
"""Generate a new random API key for Hermes callback authentication."""
|
|
return secrets.token_urlsafe(32)
|
|
|
|
|
|
def mask_api_key(key):
|
|
"""Return a masked version of the API key showing only the last 6 characters."""
|
|
if not key:
|
|
return 'Not set'
|
|
if len(key) <= 6:
|
|
return '*' * len(key)
|
|
return '*' * (len(key) - 6) + key[-6:]
|
|
|
|
|
|
def get_ntfy_config():
|
|
"""Return the effective ntfy server URL, topic, and access token from runtime settings."""
|
|
cfg = load_booth_settings()
|
|
return {
|
|
'server': cfg.get('ntfy_server', ''),
|
|
'topic': cfg.get('ntfy_topic', ''),
|
|
'token': decrypt_value(cfg.get('ntfy_token', '')) or '',
|
|
}
|
|
|
|
|
|
def send_ntfy(message, title='Theme Song Booth', priority='default', tags='bell'):
|
|
"""Send a push notification to the configured ntfy topic, if configured."""
|
|
ntfy = get_ntfy_config()
|
|
server = ntfy.get('server', '').rstrip('/')
|
|
topic = ntfy.get('topic', '').strip()
|
|
if not server or not topic:
|
|
return False
|
|
|
|
url = f"{server}/{topic}"
|
|
headers = {
|
|
'Title': title,
|
|
'Priority': priority,
|
|
'Tags': tags,
|
|
}
|
|
token = ntfy.get('token', '')
|
|
if token:
|
|
headers['Authorization'] = f'Bearer {token}'
|
|
try:
|
|
resp = requests.post(url, data=message.encode('utf-8'), headers=headers, timeout=10)
|
|
return resp.status_code in (200, 202)
|
|
except Exception:
|
|
return False
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Signed callback helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def sign_prompt_callback(rid, expires_at=None):
|
|
"""
|
|
Create a signed callback token for a specific request ID.
|
|
The signature is HMAC-SHA256 over "rid:expires_at" using APP_SECRET_KEY.
|
|
"""
|
|
secret = current_app.config['SECRET_KEY'].encode()
|
|
if expires_at is None:
|
|
expires_at = int(time.time()) + get_callback_expiry_hours() * 3600
|
|
payload = f"{rid}:{expires_at}"
|
|
sig = hmac.new(secret, payload.encode(), hashlib.sha256).hexdigest()[:16]
|
|
return f"{rid}:{expires_at}:{sig}"
|
|
|
|
|
|
def verify_prompt_callback(token):
|
|
"""Verify a signed callback token. Returns (rid, ok) tuple."""
|
|
if not token:
|
|
return None, False
|
|
parts = token.split(':')
|
|
if len(parts) != 3:
|
|
return None, False
|
|
try:
|
|
rid = int(parts[0])
|
|
expires_at = int(parts[1])
|
|
except ValueError:
|
|
return None, False
|
|
if int(time.time()) > expires_at:
|
|
return None, False
|
|
expected = sign_prompt_callback(rid, expires_at)
|
|
if not hmac.compare_digest(expected, token):
|
|
return None, False
|
|
return rid, True
|
|
|
|
|
|
def build_prompt_callback_url(rid):
|
|
"""Build the full callback URL an operator pastes into Hermes for a request."""
|
|
from flask import url_for
|
|
token = sign_prompt_callback(rid)
|
|
return f"{current_app.config['PUBLIC_BASE_URL']}/api/prompt/{rid}?token={token}"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Email helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def send_email(to, subject, body, attachments=None, inline_images=None):
|
|
"""Send an email using the configured or runtime SMTP settings."""
|
|
cfg = get_email_config()
|
|
if not cfg['SMTP_PASS']:
|
|
raise RuntimeError('SMTP password is not configured')
|
|
|
|
msg = EmailMessage()
|
|
msg['From'] = cfg['SMTP_FROM']
|
|
msg['To'] = to
|
|
msg['Subject'] = subject
|
|
msg.set_content(body)
|
|
|
|
html_body = body.replace('\n', '<br>\n')
|
|
if inline_images:
|
|
for _, cid in inline_images:
|
|
html_body += f'<br><img src="cid:{cid}" alt="Dionysis Media" style="max-width:200px;margin-top:1rem;"/>'
|
|
html_body += f'<br><br><hr style="border:none;border-top:1px solid #ddd;"/><p style="font-size:0.9rem;color:#555;">Dionysis Media: stories, sound, and a little divine chaos — <a href="https://dionysismedia.ca/">https://dionysismedia.ca/</a></p>'
|
|
msg.add_alternative(html_body, subtype='html')
|
|
|
|
if inline_images:
|
|
for path, cid in inline_images:
|
|
with open(path, 'rb') as f:
|
|
data = f.read()
|
|
ext = Path(path).suffix.lower().lstrip('.')
|
|
subtype = ext if ext in ('png', 'jpg', 'jpeg', 'gif', 'webp') else 'png'
|
|
msg.get_payload()[1].add_related(data, maintype='image', subtype=subtype, cid=f'<{cid}>')
|
|
|
|
if attachments:
|
|
for path, name in attachments:
|
|
with open(path, 'rb') as f:
|
|
data = f.read()
|
|
msg.add_attachment(data, maintype='audio', subtype='mpeg', filename=name)
|
|
|
|
with smtplib.SMTP_SSL(cfg['SMTP_HOST'], cfg['SMTP_PORT'], context=ssl.create_default_context()) as server:
|
|
server.login(cfg['SMTP_USER'], cfg['SMTP_PASS'])
|
|
server.send_message(msg)
|
|
|
|
|
|
def build_signature_images():
|
|
"""Return inline image tuple list for static/DM-Logo_email.png."""
|
|
logo_path = Path(current_app.root_path) / 'static' / 'DM-Logo_email.png'
|
|
if not logo_path.exists():
|
|
return []
|
|
return [(str(logo_path), 'dm-logo')]
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Booth state
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def get_booth_open():
|
|
"""Return True if the booth is currently marked as open in runtime settings."""
|
|
cfg = load_booth_settings()
|
|
return cfg.get('booth_open', True)
|
|
|
|
|
|
def get_musicgpt_autopoll_enabled():
|
|
"""Return True if automatic MusicGPT polling is enabled (default True)."""
|
|
cfg = load_booth_settings()
|
|
return cfg.get('musicgpt_autopoll', True)
|
|
|
|
|
|
def set_musicgpt_autopoll_enabled(enabled):
|
|
"""Persist the automatic MusicGPT polling toggle."""
|
|
cfg = load_booth_settings()
|
|
cfg['musicgpt_autopoll'] = bool(enabled)
|
|
save_booth_settings(cfg)
|
|
|
|
MUSICGPT_API_BASE = "https://api.musicgpt.com/api/public"
|
|
|
|
|
|
def _musicgpt_headers():
|
|
"""Return authorization headers for MusicGPT API calls."""
|
|
return {
|
|
"Authorization": get_musicgpt_api_key(),
|
|
"Content-Type": "application/json"
|
|
}
|
|
|
|
|
|
def musicgpt_generate_request(rid, title, music_style, lyrics, gender=None, model=None):
|
|
"""
|
|
Queue a Music AI v2 generation request.
|
|
Returns (task_id, conversion_id_1, conversion_id_2, credit_estimate, error_message).
|
|
"""
|
|
url = f"{MUSICGPT_API_BASE}/v2/MusicAI"
|
|
if model is None:
|
|
model = get_musicgpt_default_model()
|
|
payload = {
|
|
"title": title,
|
|
"music_style": music_style,
|
|
"lyrics": lyrics,
|
|
"make_instrumental": False,
|
|
"model": model,
|
|
"webhook_url": build_musicgpt_webhook_url(),
|
|
}
|
|
if gender and gender.lower() in ("male", "female", "neutral"):
|
|
payload["gender"] = gender.lower()
|
|
try:
|
|
resp = requests.post(url, json=payload, headers=_musicgpt_headers(), timeout=30)
|
|
data = resp.json()
|
|
if resp.status_code == 200 and data.get("success"):
|
|
return (
|
|
data.get("task_id"),
|
|
data.get("conversion_id_1"),
|
|
data.get("conversion_id_2"),
|
|
data.get("credit_estimate"),
|
|
None,
|
|
)
|
|
return None, None, None, None, data.get("error") or f"HTTP {resp.status_code}"
|
|
except Exception as e:
|
|
return None, None, None, None, str(e)
|
|
|
|
|
|
def musicgpt_poll_status(task_id):
|
|
"""
|
|
Poll the MusicGPT API for a generation task status.
|
|
Returns a dict with keys: status, message, conversion, or error.
|
|
"""
|
|
url = f"{MUSICGPT_API_BASE}/v1/byId"
|
|
params = {"conversionType": "MUSIC_AI", "task_id": task_id}
|
|
try:
|
|
resp = requests.get(url, headers={"Authorization": get_musicgpt_api_key()}, params=params, timeout=20)
|
|
data = resp.json()
|
|
if resp.status_code == 200:
|
|
return data
|
|
return {"error": data.get("error") or f"HTTP {resp.status_code}"}
|
|
except Exception as e:
|
|
return {"error": str(e)}
|
|
|
|
|
|
def musicgpt_queue_stems(rid, audio_url, stems=None):
|
|
"""
|
|
Queue a stem extraction job for a generated audio URL.
|
|
Returns (task_id, conversion_id, credit_estimate, error_message).
|
|
"""
|
|
url = f"{MUSICGPT_API_BASE}/v2/Extraction"
|
|
if stems is None:
|
|
stems = ["vocals", "instrumental"]
|
|
payload = {
|
|
"audio_url": audio_url,
|
|
"stems": json.dumps(stems),
|
|
"webhook_url": build_musicgpt_webhook_url(),
|
|
}
|
|
try:
|
|
resp = requests.post(url, data=payload, headers={"Authorization": get_musicgpt_api_key()}, timeout=30)
|
|
data = resp.json()
|
|
if resp.status_code == 200 and data.get("success"):
|
|
return data.get("task_id"), data.get("conversion_id"), data.get("credit_estimate"), None
|
|
return None, None, None, data.get("error") or f"HTTP {resp.status_code}"
|
|
except Exception as e:
|
|
return None, None, None, str(e)
|
|
|
|
|
|
def _download_file(url, dest):
|
|
"""Download a file from url to dest. Returns True on success."""
|
|
try:
|
|
with requests.get(url, stream=True, timeout=120) as r:
|
|
r.raise_for_status()
|
|
with open(dest, "wb") as f:
|
|
for chunk in r.iter_content(chunk_size=8192):
|
|
f.write(chunk)
|
|
return True
|
|
except Exception:
|
|
return False
|
|
|
|
|
|
def download_musicgpt_outputs(req, data, version=None):
|
|
"""
|
|
Download the MP3 and WAV outputs from a completed MusicGPT webhook/poll payload.
|
|
Updates the request row with local paths and returns a dict of saved paths.
|
|
`data` is the conversion dict from the API. For per-conversion webhooks,
|
|
pass `version='A' or 'B'`. For combined payloads, version is auto-detected
|
|
from conversion_path_1/2 keys.
|
|
"""
|
|
request_id = req["id"]
|
|
upload_dir = upload_path(request_id)
|
|
song_title = req.get("suno_title") or req.get("title")
|
|
saved = {}
|
|
|
|
# Map version label to field names.
|
|
def _fields(v):
|
|
return ("song_a_path", "song_a_wav_path") if v == "A" else ("song_b_path", "song_b_wav_path")
|
|
|
|
# Determine which version(s) are present in the payload.
|
|
versions = []
|
|
if version in ("A", "B"):
|
|
versions = [version]
|
|
elif data.get("conversion_path_1") or data.get("conversion_path_wav_1"):
|
|
if data.get("conversion_path_2") or data.get("conversion_path_wav_2"):
|
|
versions = ["A", "B"]
|
|
else:
|
|
versions = ["A"]
|
|
elif data.get("conversion_path_2") or data.get("conversion_path_wav_2"):
|
|
versions = ["B"]
|
|
elif data.get("conversion_path") or data.get("conversion_path_wav"):
|
|
# Per-conversion webhook without explicit version: use the requested version or A.
|
|
versions = [version if version in ("A", "B") else "A"]
|
|
|
|
for v in versions:
|
|
field_mp3, field_wav = _fields(v)
|
|
url_key = "conversion_path_1" if v == "A" else "conversion_path_2"
|
|
wav_key = "conversion_path_wav_1" if v == "A" else "conversion_path_wav_2"
|
|
# Also support per-conversion webhook keys without _1/_2 suffix.
|
|
if not (data.get(url_key) or data.get(wav_key)):
|
|
if v == "A" or version == "A":
|
|
mp3_url = data.get("conversion_path")
|
|
wav_url = data.get("conversion_path_wav")
|
|
else:
|
|
mp3_url = data.get("conversion_path")
|
|
wav_url = data.get("conversion_path_wav")
|
|
else:
|
|
mp3_url = data.get(url_key)
|
|
wav_url = data.get(wav_key)
|
|
|
|
if mp3_url:
|
|
ext = Path(mp3_url).suffix or ".mp3"
|
|
dest = upload_dir / f"{v}{ext}"
|
|
if _download_file(mp3_url, dest):
|
|
apply_mp3_tags(str(dest), song_title)
|
|
saved[field_mp3] = str(dest)
|
|
if wav_url:
|
|
dest_wav = upload_dir / f"{v}.wav"
|
|
if _download_file(wav_url, dest_wav):
|
|
saved[field_wav] = str(dest_wav)
|
|
|
|
# Album cover
|
|
cover_url = data.get("album_cover_path")
|
|
if cover_url:
|
|
saved["album_cover_url"] = cover_url
|
|
|
|
if saved:
|
|
update_request(request_id, **saved)
|
|
# If both song versions have been downloaded, advance request status.
|
|
from models import get_request_by_id
|
|
req = get_request_by_id(request_id)
|
|
if req and req.get('song_a_path') and req.get('song_b_path') and req.get('status') not in ('songs_uploaded', 'delivered', 'cancelled'):
|
|
update_request(request_id, status='songs_uploaded')
|
|
return saved
|
|
|
|
|
|
def format_musicgpt_cost(cost):
|
|
"""Return a human-readable cost string in USD."""
|
|
if cost is None:
|
|
return "—"
|
|
return f"${float(cost):.4f} USD"
|
|
|
|
|
|
def download_album_cover(rid, cover_url):
|
|
"""Download a MusicGPT album cover to a temp file and return its path."""
|
|
if not cover_url:
|
|
return None
|
|
try:
|
|
import tempfile
|
|
resp = requests.get(cover_url, timeout=30)
|
|
if resp.status_code == 200:
|
|
ext = Path(cover_url).suffix or '.jpg'
|
|
tmp = Path(tempfile.gettempdir()) / f"cover_{rid}{ext}"
|
|
tmp.write_bytes(resp.content)
|
|
return str(tmp)
|
|
except Exception:
|
|
pass
|
|
return None
|
|
|
|
|
|
def get_musicgpt_cost_totals():
|
|
"""Return aggregate MusicGPT cost metrics across all requests."""
|
|
from models import get_db
|
|
db = get_db()
|
|
rows = db.execute(
|
|
"""
|
|
SELECT
|
|
COUNT(*) AS song_count,
|
|
COALESCE(SUM(musicgpt_cost), 0) AS total_cost,
|
|
COALESCE(SUM(
|
|
CASE WHEN song_a_path IS NOT NULL AND song_b_path IS NOT NULL THEN 1 ELSE 0
|
|
END), 0) AS completed_pairs,
|
|
COALESCE(SUM(stems_cost), 0) AS stems_total
|
|
FROM requests
|
|
WHERE musicgpt_cost IS NOT NULL
|
|
"""
|
|
).fetchone()
|
|
return {
|
|
"song_count": rows["song_count"] or 0,
|
|
"completed_pairs": rows["completed_pairs"] or 0,
|
|
"total_cost": rows["total_cost"] or 0,
|
|
"stems_total": rows["stems_total"] or 0,
|
|
}
|