260 lines
8.1 KiB
Python

import os
import subprocess
import tempfile
import re
import logging
import glob
from pathlib import Path
from flask import Flask, request, jsonify
from werkzeug.utils import secure_filename
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
app = Flask(__name__)
# Configuration from environment variables
OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "output")
PORT = int(os.environ.get("PORT", 5012))
API_KEY = os.environ.get("TTS_API_KEY", "")
# Issue #10: Whitelist allowed TTS commands
ALLOWED_TTS_COMMANDS = {"kokoro-tts", "/usr/local/bin/kokoro-tts", "/usr/bin/kokoro-tts"}
_tts_command = os.environ.get("TTS_COMMAND", "kokoro-tts")
if _tts_command not in ALLOWED_TTS_COMMANDS:
logger.warning(
f"TTS_COMMAND '{_tts_command}' not in whitelist, using default 'kokoro-tts'"
)
TTS_COMMAND = "kokoro-tts"
else:
TTS_COMMAND = _tts_command
# Issue #12: Maximum number of output files to keep per title
MAX_OUTPUT_FILES_PER_TITLE = 5
# Ensure output directory exists
os.makedirs(OUTPUT_DIR, exist_ok=True)
# Available voices
VOICES = ["bm_fable", "bm_lewis", "bm_george"]
def parse_speaker_text(text):
"""Parse text into speaker: complete lines pairs for full conversation"""
lines = [line.strip() for line in text.strip().split("\n") if line.strip()]
speakers = {}
current_speaker = None
speaker_content = []
tracks = []
for line in lines:
match = re.match(r"^(.*?): (.*)$", line)
if match:
if current_speaker and speaker_content:
speakers[current_speaker] = "\n".join(speaker_content)
current_speaker = match.group(1).strip()
speaker_content = [match.group(2).strip()]
tracks += [(current_speaker, speaker_content[0])]
else:
if current_speaker and line.strip():
speaker_content.append(line.strip())
if current_speaker and speaker_content:
speakers[current_speaker] = "\n".join(speaker_content)
return (speakers, tracks)
def assign_voices_to_speakers(speakers):
"""Assign voices to speakers (same voice for each speaker throughout conversation)"""
speaker_voice_map = {}
used_voices = set()
for speaker in speakers.keys():
if not speaker_voice_map.get(speaker):
available_voices = [voice for voice in VOICES if voice not in used_voices]
if available_voices:
selected_voice = available_voices[0]
else:
selected_voice = VOICES[0]
speaker_voice_map[speaker] = selected_voice
used_voices.add(selected_voice)
return speaker_voice_map
def clean_old_outputs(title):
"""Issue #12: Remove old output files for a given title, keeping only the newest."""
pattern = os.path.join(OUTPUT_DIR, f"{title}_*.wav")
files = sorted(glob.glob(pattern), key=os.path.getmtime, reverse=True)
for old_file in files[MAX_OUTPUT_FILES_PER_TITLE:]:
try:
os.unlink(old_file)
logger.info(f"Removed old output file: {old_file}")
except OSError as e:
logger.warning(f"Failed to remove old output file {old_file}: {e}")
def generate_wav_files(speakers, speaker_voice_map, title, tracks):
"""Generate individual WAV files for each speaker"""
wav_files = []
turn = 0
for speaker, text in tracks:
safe_speaker = secure_filename(speaker)
filepath = os.path.join(OUTPUT_DIR, f"{title}_{turn}.wav")
turn += 1
temp_input_path = None
try:
# Issue #11: Create temporary input file with proper cleanup
with tempfile.NamedTemporaryFile(
mode="w", suffix=".txt", delete=False
) as temp_file:
temp_file.write(text)
temp_input_path = temp_file.name
cmd = [
TTS_COMMAND,
temp_input_path,
filepath,
"--voice",
speaker_voice_map[speaker],
]
subprocess.run(cmd, check=True, capture_output=True)
wav_files.append(filepath)
except Exception as e:
logger.error(f"Error generating audio for {speaker}: {e}")
return None
finally:
# Issue #11: Always clean up temp file, using specific exception type
if temp_input_path and os.path.exists(temp_input_path):
try:
os.unlink(temp_input_path)
except OSError as cleanup_err:
logger.warning(
f"Failed to clean up temp file {temp_input_path}: {cleanup_err}"
)
return wav_files
def merge_wav_files(wav_files, final_output_path):
"""Merge multiple WAV files into a single file using sox"""
try:
if not wav_files:
return None
if len(wav_files) == 1:
import shutil
shutil.copy2(wav_files[0], final_output_path)
return final_output_path
cmd = ["sox"] + wav_files + [final_output_path]
subprocess.run(cmd, check=True, capture_output=True)
# Clean up intermediate files
for wav_file in wav_files:
try:
os.unlink(wav_file)
except OSError as e:
logger.warning(f"Failed to remove intermediate file {wav_file}: {e}")
return final_output_path
except Exception as e:
logger.error(f"Error merging WAV files: {e}")
if wav_files and len(wav_files) > 0:
import shutil
try:
shutil.copy2(wav_files[-1], final_output_path)
return final_output_path
except OSError as fallback_err:
logger.error(
f"Fallback copy failed for {wav_files[-1]}: {fallback_err}"
)
return None
def check_api_key():
"""Issue #9: Validate API key from Authorization header."""
if not API_KEY:
return True # No key configured = skip auth (dev mode)
auth_header = request.headers.get("Authorization", "")
if not auth_header.startswith("Bearer "):
return False
provided_key = auth_header[len("Bearer "):]
return provided_key == API_KEY
@app.route("/tts", methods=["POST"])
def text_to_speech():
"""Endpoint to convert full conversation to audio podcast"""
# Issue #9: API key authentication
if not check_api_key():
return jsonify({"error": "Unauthorized. Provide valid API key."}), 401
data = request.get_json()
if not data:
return jsonify({"error": "No JSON data provided"}), 400
text = data.get("text")
title = data.get("title")
if not text or not title:
return jsonify({"error": "Both 'text' and 'title' fields are required"}), 400
# Sanitize title to prevent path traversal
safe_title = secure_filename(title)
if not safe_title:
return jsonify({"error": "Invalid title format"}), 400
(speakers, tracks) = parse_speaker_text(text)
if not speakers:
return jsonify({"error": "No speaker lines found in text"}), 400
speaker_voice_map = assign_voices_to_speakers(speakers)
# Issue #12: Clean old outputs before generating new ones
clean_old_outputs(safe_title)
wav_files = generate_wav_files(speakers, speaker_voice_map, safe_title, tracks)
if not wav_files:
return jsonify({"error": "Failed to generate audio files"}), 500
final_wav_path = os.path.join(OUTPUT_DIR, f"{safe_title}.wav")
merged_file = merge_wav_files(wav_files, final_wav_path)
if not merged_file:
return jsonify({"error": "Failed to merge audio files"}), 500
return jsonify(
{
"file_path": final_wav_path,
"message": "Full conversation podcast generated successfully",
"speakers": list(speakers.keys()),
"total_speakers": len(speakers),
"total_lines": len(text.split("\n")),
}
)
@app.route("/health", methods=["GET"])
def health_check():
"""Health check endpoint"""
return jsonify({"status": "healthy"})
if __name__ == "__main__":
app.run(host="0.0.0.0", port=PORT, debug=False)