#!/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()