130 lines
3.9 KiB
Python
130 lines
3.9 KiB
Python
"""Pipeline orchestrator: runs all three stages and schedules repeated execution."""
|
|
|
|
import asyncio
|
|
import os
|
|
import signal
|
|
import sys
|
|
import threading
|
|
from loguru import logger
|
|
|
|
# --- Configuration ---
|
|
from server.config import resolve_ollama_url
|
|
|
|
OLLAMA_URL = resolve_ollama_url(
|
|
os.environ.get("OLLAMA_URL", "http://127.0.0.1:11434/api/generate")
|
|
)
|
|
MODEL_NAME = os.environ.get("MODEL_NAME", "gemma4:e2b")
|
|
PORT = int(os.environ.get("PORT", "8080"))
|
|
POLL_INTERVAL = int(os.environ.get("POLL_INTERVAL", "21600")) # 6 hours
|
|
DATA_DIR = os.environ.get("DATA_DIR", "/app/data")
|
|
TTS_VOICE = os.environ.get("TTS_VOICE", "af_heart")
|
|
|
|
# Ensure all stages read the same env vars
|
|
os.environ["OLLAMA_URL"] = OLLAMA_URL
|
|
os.environ["MODEL_NAME"] = MODEL_NAME
|
|
os.environ["DATA_DIR"] = DATA_DIR
|
|
os.environ["TTS_VOICE"] = TTS_VOICE
|
|
|
|
# Import pipeline stages
|
|
from main import main as stage1_scrape
|
|
from reader import main as stage2_script
|
|
from tts_generator import main as stage3_tts
|
|
|
|
# Start web server in background thread
|
|
from server.app import run_server as start_server
|
|
|
|
# --- Logging ---
|
|
logger.add("pipeline.log", rotation="5 MB", retention="7 days")
|
|
|
|
|
|
async def run_pipeline() -> bool:
|
|
"""Run all three pipeline stages sequentially. Returns True on success."""
|
|
logger.info("=" * 60)
|
|
logger.info("Stage 1: Scraping Guardian articles...")
|
|
try:
|
|
await stage1_scrape()
|
|
except Exception:
|
|
logger.exception("Stage 1 (scrape) failed.")
|
|
return False
|
|
|
|
output_file = os.path.join(DATA_DIR, "output.json")
|
|
if not os.path.exists(output_file):
|
|
logger.error("Stage 1 produced no output. Skipping stages 2 and 3.")
|
|
return False
|
|
|
|
logger.info("Stage 2: Generating newscast scripts...")
|
|
try:
|
|
await stage2_script()
|
|
except Exception:
|
|
logger.exception("Stage 2 (script generation) failed.")
|
|
return False
|
|
|
|
scripts_file = os.path.join(DATA_DIR, "scripts.json")
|
|
if not os.path.exists(scripts_file):
|
|
logger.error("Stage 2 produced no output. Skipping stage 3.")
|
|
return False
|
|
|
|
logger.info("Stage 3: Generating audio files...")
|
|
try:
|
|
await stage3_tts()
|
|
except Exception:
|
|
logger.exception("Stage 3 (TTS) failed.")
|
|
return False
|
|
|
|
logger.info(f"Pipeline complete! Audio files available in {DATA_DIR}/audio/")
|
|
return True
|
|
|
|
|
|
async def main():
|
|
logger.info(f"Guardian Daily Newscast Pipeline starting")
|
|
logger.info(f" DATA_DIR = {DATA_DIR}")
|
|
logger.info(f" OLLAMA_URL = {OLLAMA_URL}")
|
|
logger.info(f" MODEL = {MODEL_NAME}")
|
|
logger.info(f" VOICE = {TTS_VOICE}")
|
|
logger.info(f" PORT = {PORT}")
|
|
logger.info(f" POLL_INTERVAL = {POLL_INTERVAL}s ({POLL_INTERVAL // 3600}h)")
|
|
|
|
# Signal handling for graceful shutdown
|
|
loop = asyncio.get_running_loop()
|
|
stop = loop.create_future()
|
|
|
|
def _signal_handler():
|
|
if not stop.done():
|
|
stop.set_result(None)
|
|
|
|
for sig in (signal.SIGINT, signal.SIGTERM):
|
|
loop.add_signal_handler(sig, _signal_handler)
|
|
|
|
# Start web server in background thread
|
|
server_thread = threading.Thread(
|
|
target=start_server, args=(PORT, DATA_DIR), daemon=True
|
|
)
|
|
server_thread.start()
|
|
|
|
run_count = 0
|
|
while True:
|
|
run_count += 1
|
|
logger.info(f"--- Run #{run_count} ---")
|
|
try:
|
|
success = await run_pipeline()
|
|
except Exception:
|
|
logger.exception("Pipeline error.")
|
|
success = False
|
|
|
|
if success:
|
|
logger.info(f"Next run in {POLL_INTERVAL // 3600} hours ({POLL_INTERVAL}s)...")
|
|
else:
|
|
logger.warning("Pipeline had errors. Waiting before retry...")
|
|
|
|
try:
|
|
await asyncio.wait_for(stop, timeout=POLL_INTERVAL)
|
|
break
|
|
except asyncio.TimeoutError:
|
|
pass
|
|
|
|
logger.info("Shutting down.")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
asyncio.run(main())
|