CoolFace
Apppublic

gribsons/quant-tradingbot

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
api_server.py266 linesDownload Raw Back to root
1"""2api_server.py — FastAPI backend for QuantBot.3 4Exposes the backtesting engine over HTTP so the Next.js frontend5(or any HTTP client) can run backtests without touching Python directly.6 7Deploy to Railway:8    Procfile: web: uvicorn api_server:app --host 0.0.0.0 --port $PORT9"""10 11from __future__ import annotations12 13import logging14import math15import sys16import os17 18import numpy as np19import pandas as pd20from fastapi import FastAPI, HTTPException21from fastapi.middleware.cors import CORSMiddleware22from pydantic import BaseModel, Field23 24# Local imports (this file lives at the quant_tradingbot root)25from config import (26    AnalyticsConfig,27    BrokerConfig,28    EnsembleConfig,29    MACrossoverConfig,30    MomentumConfig,31    QuantBotConfig,32    RiskConfig,33    RSIReversionConfig,34)35from data.fetcher import DataFetcher36from data.processor import DataProcessor37from backtest.engine import BacktestEngine38from risk.manager import RiskManager39from strategies import STRATEGY_REGISTRY, get_strategy40 41logging.basicConfig(level=logging.INFO)42logger = logging.getLogger(__name__)43 44app = FastAPI(title="QuantBot API", version="1.0.0")45 46# ---------------------------------------------------------------------------47# CORS — allow the Vercel frontend (and localhost dev)48# ---------------------------------------------------------------------------49 50ALLOWED_ORIGINS = os.environ.get(51    "ALLOWED_ORIGINS",52    "http://localhost:3000,http://localhost:3001",53).split(",")54 55app.add_middleware(56    CORSMiddleware,57    allow_origins=ALLOWED_ORIGINS,58    allow_origin_regex=r"https://.*\.vercel\.app",59    allow_credentials=True,60    allow_methods=["*"],61    allow_headers=["*"],62)63 64 65# ---------------------------------------------------------------------------66# Request / Response models67# ---------------------------------------------------------------------------68 69class BacktestRequest(BaseModel):70    symbol: str = "SPY"71    start_date: str = "2020-01-01"72    end_date: str = "2024-12-31"73    interval: str = "1d"74    strategy: str = "ensemble"75    # Broker76    initial_cash: float = Field(100_000.0, ge=1_000)77    commission_bps: float = Field(10.0, ge=0, le=100)78    slippage_bps: float = Field(5.0, ge=0, le=100)79    # Risk80    max_position_pct: float = Field(0.20, ge=0.01, le=1.0)81    max_drawdown_pct: float = Field(0.15, ge=0.01, le=1.0)82    max_open_positions: int = Field(5, ge=1, le=20)83    position_sizing: str = "fixed_pct"84    # Ensemble-specific85    adx_trend: float = 25.086    adx_range: float = 20.087    vote_threshold: float = 0.588    min_agreement: int = 289    atr_stop: float = 2.090    atr_target: float = 4.091 92 93class EquityPoint(BaseModel):94    date: str95    equity: float96 97 98class TradeRecord(BaseModel):99    entry_date: str = ""100    exit_date: str = ""101    symbol: str = ""102    direction: str = ""103    entry_price: float = 0.0104    exit_price: float = 0.0105    shares: float = 0.0106    pnl: float = 0.0107    pnl_pct: float = 0.0108    commission: float = 0.0109    slippage: float = 0.0110    exit_reason: str = ""111 112 113class BacktestResponse(BaseModel):114    strategy_name: str115    symbol: str116    metrics: dict117    equity_curve: list[EquityPoint]118    trade_log: list[dict]119 120 121# ---------------------------------------------------------------------------122# Helpers123# ---------------------------------------------------------------------------124 125def _safe(v):126    """Make a value JSON-safe (handles inf / nan)."""127    if isinstance(v, float):128        if math.isnan(v) or math.isinf(v):129            return None130    return v131 132 133def _df_to_equity_points(ec: pd.DataFrame) -> list[dict]:134    """Convert equity curve DataFrame to list of {date, equity} dicts."""135    if ec.empty:136        return []137    result = []138    for idx, row in ec.iterrows():139        result.append({140            "date": str(idx.date()) if hasattr(idx, "date") else str(idx),141            "equity": round(float(row["equity"]), 2),142        })143    return result144 145 146def _df_to_trade_records(tl: pd.DataFrame) -> list[dict]:147    """Convert trade log DataFrame to list of dicts, JSON-safe."""148    if tl.empty:149        return []150    records = []151    for _, row in tl.iterrows():152        rec = {}153        for col in tl.columns:154            val = row[col]155            if isinstance(val, (pd.Timestamp,)):156                rec[col] = str(val.date())157            elif isinstance(val, float) and (math.isnan(val) or math.isinf(val)):158                rec[col] = None159            elif isinstance(val, (np.integer,)):160                rec[col] = int(val)161            elif isinstance(val, (np.floating,)):162                rec[col] = round(float(val), 6)163            else:164                rec[col] = val165        records.append(rec)166    return records167 168 169# ---------------------------------------------------------------------------170# Routes171# ---------------------------------------------------------------------------172 173@app.get("/")174def root():175    return {"name": "QuantBot API", "version": "1.0.0", "docs": "/docs", "health": "/health"}176 177 178@app.get("/health")179def health():180    return {"status": "ok", "strategies": list(STRATEGY_REGISTRY.keys())}181 182 183@app.get("/strategies")184def list_strategies():185    return {"strategies": list(STRATEGY_REGISTRY.keys())}186 187 188@app.post("/run", response_model=BacktestResponse)189def run_backtest(req: BacktestRequest):190    logger.info("Backtest request: %s %s %s→%s", req.strategy, req.symbol, req.start_date, req.end_date)191 192    # Build config193    cfg = QuantBotConfig(194        broker=BrokerConfig(195            commission_pct=req.commission_bps / 10_000.0,196            slippage_pct=req.slippage_bps / 10_000.0,197            initial_cash=req.initial_cash,198        ),199        risk=RiskConfig(200            max_position_pct=req.max_position_pct,201            max_drawdown_pct=req.max_drawdown_pct,202            max_open_positions=req.max_open_positions,203            position_sizing=req.position_sizing,204        ),205        ma_crossover=MACrossoverConfig(),206        rsi_reversion=RSIReversionConfig(),207        momentum=MomentumConfig(),208        ensemble=EnsembleConfig(209            adx_trend_threshold=req.adx_trend,210            adx_range_threshold=req.adx_range,211            vote_threshold=req.vote_threshold,212            min_agreement=int(req.min_agreement),213            atr_stop_multiplier=req.atr_stop,214            atr_target_multiplier=req.atr_target,215        ),216        analytics=AnalyticsConfig(),217        start_date=req.start_date,218        end_date=req.end_date,219        interval=req.interval,220    )221 222    try:223        # Fetch data224        raw = DataFetcher().fetch(req.symbol, req.start_date, req.end_date, req.interval)225    except Exception as exc:226        raise HTTPException(status_code=422, detail=f"Data fetch failed: {exc}")227 228    try:229        proc = DataProcessor().process(raw)230    except Exception as exc:231        raise HTTPException(status_code=422, detail=f"Data processing failed: {exc}")232 233    # Build strategy kwargs234    kwargs: dict = {}235    if req.strategy == "ensemble":236        kwargs = {"config": cfg.ensemble}237    elif req.strategy == "ma_crossover":238        kwargs = {"config": cfg.ma_crossover}239    elif req.strategy == "rsi_reversion":240        kwargs = {"config": cfg.rsi_reversion}241    elif req.strategy == "momentum":242        kwargs = {"config": cfg.momentum}243 244    try:245        strat = get_strategy(req.strategy, **kwargs)246    except Exception as exc:247        raise HTTPException(status_code=422, detail=f"Strategy error: {exc}")248 249    try:250        rm = RiskManager(cfg.risk, cfg.broker.initial_cash)251        result = BacktestEngine(cfg).run(strat, proc, req.symbol, rm)252    except Exception as exc:253        logger.exception("Backtest failed")254        raise HTTPException(status_code=500, detail=f"Backtest failed: {exc}")255 256    # Sanitise metrics257    safe_metrics = {k: _safe(v) for k, v in result.metrics.items()}258 259    return BacktestResponse(260        strategy_name=result.strategy_name,261        symbol=result.symbol,262        metrics=safe_metrics,263        equity_curve=_df_to_equity_points(result.equity_curve),264        trade_log=_df_to_trade_records(result.trade_log),265    )266