#!/usr/bin/env python3
"""Scheduler — auto-fetch data and retrain models every 4 hours."""
import argparse
import logging
import sys
import os
import signal
import time
from datetime import datetime

# Add project root to path
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, PROJECT_ROOT)

from apscheduler.schedulers.blocking import BlockingScheduler
from backend.config import config
from backend.data import fetch_all_tickers
from backend.features import compute_features, engineer_features, prepare_training_data
from backend.predictor import Predictor

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(message)s",
    datefmt="%Y-%m-%d %H:%M:%S",
)
logger = logging.getLogger(__name__)

running = True


def signal_handler(sig, frame):
    global running
    logger.info("Shutting down scheduler...")
    running = False


signal.signal(signal.SIGINT, signal_handler)
signal.signal(signal.SIGTERM, signal_handler)


def fetch_data():
    """Fetch latest data for all tickers."""
    logger.info("Fetching latest data for all tickers...")
    results = fetch_all_tickers()
    logger.info(f"Data fetch complete: {len(results)} tickers")
    return results


def train_ticker(ticker, df):
    """Train model for a single ticker."""
    from backend.features import compute_features, engineer_features, prepare_training_data
    
    df_features = compute_features(df)
    lookahead = config.data.get("lookahead", 5)
    df_targets = engineer_features(df_features, lookahead=lookahead)
    
    X, y_class, y_reg, feature_names = prepare_training_data(df_targets)
    
    if len(X) < 100:
        logger.warning(f"Not enough data for {ticker}: {len(X)} bars")
        return None
    
    pred = Predictor()
    
    # Try loading existing model to compare
    if pred.load(ticker):
        logger.info(f"Existing model for {ticker} (trained {pred.trained_at})")
    
    metrics = pred.train(X, y_class, y_reg, feature_names, ticker=ticker)
    pred.save(ticker)
    
    acc = metrics["classification"]["accuracy"]
    r2 = metrics["regression"]["r2"]
    logger.info(f"  {ticker}: acc={acc:.3f}, r2={r2:.3f}, n={len(X)}")
    return metrics


def train_all():
    """Train/retrain models for all tickers."""
    logger.info("Training models for all tickers...")
    cached = fetch_all_tickers()
    
    for ticker, df in cached.items():
        try:
            metrics = train_ticker(ticker, df)
            if metrics:
                logger.info(f"  ✓ {ticker} trained successfully")
            else:
                logger.warning(f"  ✗ {ticker} training skipped")
        except Exception as e:
            logger.error(f"  ✗ {ticker}: {e}")


def run_cycle():
    """Full cycle: fetch + train if configured."""
    logger.info("=" * 50)
    logger.info(f"Scheduler cycle at {datetime.now().isoformat()}")
    
    try:
        fetch_data()
        
        if config.scheduler.get("auto_retrain", True):
            train_all()
        
        logger.info("Cycle complete")
    except Exception as e:
        logger.error(f"Cycle failed: {e}", exc_info=True)


def main():
    parser = argparse.ArgumentParser(description="Stock Prediction Engine")
    parser.add_argument("--fetch-only", action="store_true", help="Just fetch data, don't train")
    parser.add_argument("--train-only", action="store_true", help="Just train, don't fetch")
    parser.add_argument("--once", action="store_true", help="Run once and exit")
    parser.add_argument("--server", action="store_true", help="Start FastAPI server")
    parser.add_argument("--scheduler", action="store_true", help="Run scheduler loop")
    parser.add_argument("--all", action="store_true", help="Fetch + train + server + scheduler")
    parser.add_argument("--ticker", type=str, help="Specific ticker to fetch/train")
    args = parser.parse_args()
    
    if args.server or args.all:
        # Start server in background thread
        import threading
        import uvicorn
        
        def run_server():
            uvicorn.run(
                "backend.app:app",
                host=config.app.get("host", "0.0.0.0"),
                port=config.app.get("port", 5080),
                log_level="info",
            )
        
        server_thread = threading.Thread(target=run_server, daemon=True)
        server_thread.start()
        logger.info(f"Server starting on {config.app.get('host')}:{config.app.get('port')}")
    
    # Initial run
    if args.fetch_only or args.train_only or args.once or args.all:
        run_cycle()
        if args.once:
            logger.info("One-shot complete, exiting")
            return
    
    if not (args.server or args.all):
        run_cycle()
        return
    
    # Scheduler loop
    if args.scheduler or args.all:
        interval = config.scheduler.get("interval_hours", 4)
        logger.info(f"Scheduler running, interval={interval}h")
        
        scheduler = BlockingScheduler()
        scheduler.add_job(
            run_cycle,
            "interval",
            hours=interval,
            id="stock_prediction_cycle",
            name="Fetch & train all tickers",
        )
        
        try:
            scheduler.start()
        except (KeyboardInterrupt, SystemExit):
            logger.info("Scheduler stopped")


if __name__ == "__main__":
    main()
