"""FastAPI application — REST API for stock prediction engine."""
import asyncio
import logging
import os
from datetime import datetime
from typing import Dict, List, Optional
import numpy as np
import pandas as pd
from fastapi import FastAPI, HTTPException, Query
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, JSONResponse
from fastapi.staticfiles import StaticFiles
from backend.config import config
from backend.data import fetch_all_tickers, fetch_ticker, load_all_cached, load_cached, get_latest_bar
from backend.features import compute_features, engineer_features, prepare_training_data
from backend.predictor import Predictor
logger = logging.getLogger(__name__)
app = FastAPI(
title="Stock Prediction Engine",
description="4H OHLCV analysis, technical indicators, and ML predictions",
version="1.0.0",
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
PROJECT_ROOT = os.path.dirname(os.path.dirname(__file__))
# --- Global model registry ---
models: Dict[str, Predictor] = {}
@app.on_event("startup")
async def startup():
"""Load cached models on startup."""
for ticker in config.data.get("tickers", []):
pred = Predictor()
if pred.load(ticker):
models[ticker] = pred
logger.info(f"Loaded model for {ticker}")
@app.get("/")
async def root():
return FileResponse(os.path.join(PROJECT_ROOT, "frontend", "index.html"))
@app.get("/health")
async def health():
return {
"status": "ok",
"timestamp": datetime.now().isoformat(),
"tickers_tracked": len(config.data.get("tickers", [])),
"models_loaded": len(models),
}
# --- Data endpoints ---
@app.get("/api/tickers")
async def list_tickers():
"""List all tracked tickers with latest price."""
result = []
for ticker in config.data.get("tickers", []):
latest = get_latest_bar(ticker)
model_info = models.get(ticker)
result.append({
"ticker": ticker,
"latest": latest,
"has_model": model_info.is_trained if model_info else False,
"model_accuracy": model_info.train_metrics.get("classification", {}).get("accuracy") if model_info and model_info.is_trained else None,
})
return {"tickers": result}
@app.get("/api/data/{ticker}")
async def get_ticker_data(
ticker: str,
limit: int = Query(500, ge=1, le=10000),
):
"""Get OHLCV + indicators for a ticker (last N bars)."""
df = load_cached(ticker)
if df is None:
df = fetch_ticker(ticker)
if df.empty:
raise HTTPException(404, f"No data for {ticker}")
df = compute_features(df)
# Return last N bars
df = df.tail(limit).reset_index()
# Convert to list of dicts
records = []
for _, row in df.iterrows():
rec = {}
for col in df.columns:
val = row[col]
if isinstance(val, (np.floating, float)):
rec[col] = round(float(val), 6) if not np.isnan(val) else None
elif isinstance(val, (np.integer, int)):
rec[col] = int(val)
elif isinstance(val, pd.Timestamp):
rec[col] = str(val)
else:
rec[col] = val
records.append(rec)
return {
"ticker": ticker,
"count": len(records),
"data": records,
}
# --- Prediction endpoints ---
@app.post("/api/train/{ticker}")
async def train_ticker(ticker: str):
"""Train (retrain) model for a ticker."""
df = load_cached(ticker)
if df is None or df.empty:
df = fetch_ticker(ticker)
if df.empty:
raise HTTPException(404, f"No data for {ticker}")
# Compute features + targets
df = compute_features(df)
lookahead = config.data.get("lookahead", 5)
df = engineer_features(df, lookahead=lookahead)
X, y_class, y_reg, feature_names = prepare_training_data(df)
if len(X) < 100:
raise HTTPException(422, f"Not enough data for training: {len(X)} bars")
pred = Predictor()
metrics = pred.train(X, y_class, y_reg, feature_names, ticker=ticker)
pred.save(ticker)
models[ticker] = pred
return {
"ticker": ticker,
"status": "trained",
"metrics": metrics,
"trained_at": pred.trained_at,
}
@app.post("/api/train_all")
async def train_all():
"""Train models for all tickers."""
results = []
for ticker in config.data.get("tickers", []):
try:
result = await train_ticker(ticker)
results.append({"ticker": ticker, "status": "ok", "metrics": result.get("metrics")})
except Exception as e:
results.append({"ticker": ticker, "status": "error", "error": str(e)})
return {"status": "complete", "results": results}
@app.get("/api/predict/{ticker}")
async def predict_ticker(ticker: str):
"""Get prediction for the latest bar of a ticker."""
pred = models.get(ticker)
if not pred or not pred.is_trained:
raise HTTPException(404, f"No trained model for {ticker}. Train first.")
df = load_cached(ticker)
if df is None or df.empty:
raise HTTPException(404, f"No cached data for {ticker}")
# Get latest bar features
df = compute_features(df)
features = [c for c in df.columns if c not in {"Date", "future_close", "target_pct", "target_direction"}]
# Latest complete bar (exclude last row which may be incomplete)
latest = df.iloc[-1]
X = latest[features].values.reshape(1, -1).astype(np.float32)
prediction = pred.predict(X)
prediction["ticker"] = ticker
prediction["date"] = str(df.index[-1])
prediction["close"] = round(float(latest["Close"]), 2)
# Add model info
prediction["model_accuracy"] = pred.train_metrics.get("classification", {}).get("accuracy")
prediction["model_trained"] = pred.trained_at
return prediction
@app.get("/api/model/{ticker}")
async def model_info(ticker: str):
"""Get model training metrics and feature importance."""
pred = models.get(ticker)
if not pred or not pred.is_trained:
raise HTTPException(404, f"No trained model for {ticker}")
return pred.info()
# --- Alert endpoints ---
@app.get("/api/alerts")
async def get_alerts():
"""Check all tickers for alert conditions."""
alerts = []
ac = config.alerts
for ticker in config.data.get("tickers", []):
df = load_cached(ticker)
if df is None or df.empty:
continue
df = compute_features(df)
latest = df.iloc[-1]
prev = df.iloc[-2] if len(df) > 1 else latest
# Price change alert
price_change = (latest["Close"] - prev["Close"]) / prev["Close"] * 100
if abs(price_change) >= ac.get("price_change_pct", 5.0):
alerts.append({
"ticker": ticker,
"type": "price_change",
"severity": "high" if abs(price_change) > 10 else "medium",
"message": f"{ticker}: {price_change:+.2f}% price change",
"value": round(price_change, 2),
"timestamp": str(df.index[-1]),
})
# RSI alerts
rsi = latest.get("RSI")
if rsi is not None and not np.isnan(rsi):
if rsi >= ac.get("rsi_overbought", 70):
alerts.append({
"ticker": ticker,
"type": "rsi_overbought",
"severity": "medium",
"message": f"{ticker}: RSI overbought at {rsi:.1f}",
"value": round(float(rsi), 1),
"timestamp": str(df.index[-1]),
})
elif rsi <= ac.get("rsi_oversold", 30):
alerts.append({
"ticker": ticker,
"type": "rsi_oversold",
"severity": "medium",
"message": f"{ticker}: RSI oversold at {rsi:.1f}",
"value": round(float(rsi), 1),
"timestamp": str(df.index[-1]),
})
# MACD crossover
if ac.get("macd_crossover"):
macd_hist = latest.get("MACD_histogram")
macd_hist_prev = prev.get("MACD_histogram")
if macd_hist is not None and macd_hist_prev is not None:
if not np.isnan(macd_hist) and not np.isnan(macd_hist_prev):
if macd_hist_prev < 0 and macd_hist > 0:
alerts.append({
"ticker": ticker,
"type": "macd_bullish",
"severity": "low",
"message": f"{ticker}: MACD bullish crossover",
"value": round(float(macd_hist), 4),
"timestamp": str(df.index[-1]),
})
elif macd_hist_prev > 0 and macd_hist < 0:
alerts.append({
"ticker": ticker,
"type": "macd_bearish",
"severity": "low",
"message": f"{ticker}: MACD bearish crossover",
"value": round(float(macd_hist), 4),
"timestamp": str(df.index[-1]),
})
# High-confidence prediction alert
pred = models.get(ticker)
if pred and pred.is_trained:
try:
features = [c for c in df.columns if c not in {"Date", "future_close", "target_pct", "target_direction"}]
X = df.iloc[-1][features].values.reshape(1, -1).astype(np.float32)
p = pred.predict(X)
if p["confidence"] >= ac.get("prediction_confidence", 0.65) and p["direction"] != "neutral":
alerts.append({
"ticker": ticker,
"type": "prediction",
"severity": "high" if p["confidence"] > 0.8 else "medium",
"message": f"{ticker}: {p['direction'].upper()} prediction ({p['confidence']:.1%} conf, {p['pct_change']:+.2f}%)",
"direction": p["direction"],
"confidence": p["confidence"],
"pct_change": p["pct_change"],
"timestamp": str(df.index[-1]),
})
except Exception:
pass
return {"alerts": alerts, "count": len(alerts)}
# --- Config endpoint ---
@app.get("/api/config")
async def get_config():
"""Get current config (sans secrets)."""
return {
"app": config.app,
"data": config.data,
"features": config.features,
"model": config.model,
"alerts": config.alerts,
}
@app.get("/api/dashboard")
async def dashboard_data():
"""Aggregate dashboard data for all tickers."""
tickers = config.data.get("tickers", [])
result = []
for ticker in tickers:
latest = get_latest_bar(ticker)
pred = models.get(ticker)
prediction = None
if pred and pred.is_trained:
try:
df = load_cached(ticker)
if df is not None and not df.empty:
df = compute_features(df)
features = [c for c in df.columns if c not in {"Date", "future_close", "target_pct", "target_direction"}]
X = df.iloc[-1][features].values.reshape(1, -1).astype(np.float32)
prediction = pred.predict(X)
prediction["ticker"] = ticker
except Exception:
pass
result.append({
"ticker": ticker,
"latest": latest,
"prediction": prediction,
"has_model": pred.is_trained if pred else False,
})
return {"tickers": result, "updated_at": datetime.now().isoformat()}
if __name__ == "__main__":
import uvicorn
uvicorn.run("backend.app:app", host=config.app.get("host"), port=config.app.get("port"), reload=True)