"""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)