CoolFace
Apppublic

thinkingEverytime/QuantOracle

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
publish_eod.py572 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""Build + publish EOD screener artifacts (features + model) for a universe.3 4Goal: produce a stable, daily "as-of close" snapshot for Streamlit Cloud.5 6Outputs (local):7  data/features.parquet8  data/models/ridge_h{horizon}/<version>/{model.npz,meta.json}9  data/models/ridge_h{horizon}/LATEST10  data/eod_latest.json11 12If --upload is set, uploads to Supabase Storage public bucket under:13  eod/<universe>/features.parquet14  eod/<universe>/models/...15  eod/<universe>/latest.json16"""17 18# ruff: noqa: E402  (sys.path bootstrap must run before local imports)19 20from __future__ import annotations21 22import argparse23import json24import os25import re26import sys27from dataclasses import dataclass28from datetime import datetime, timedelta, timezone29from pathlib import Path30from typing import Any, Iterable31from urllib.parse import quote32 33import duckdb34import numpy as np35import pandas as pd36import requests37from dotenv import load_dotenv38 39ROOT = Path(__file__).resolve().parents[1]40if str(ROOT) not in sys.path:41    sys.path.insert(0, str(ROOT))42 43load_dotenv()44 45from quant.features import build_features, build_targets46from quant.registry import (47    data_root,48    model_version_dir,49    save_meta,50    version_id,51    write_latest,52)53from quant.ridge import fit_ridge, predict, zscore_apply, zscore_fit54from scripts.groww_api import GrowwAuth, get_access_token, get_candles_range55from scripts.supabase_storage import from_env56 57 58@dataclass(frozen=True)59class Provider:60    name: str61 62    def history(self, sym: str, *, days: int) -> pd.DataFrame:63        raise NotImplementedError64 65 66class YFinance(Provider):67    def __init__(self):68        super().__init__(name="yfinance")69 70    def history(self, sym: str, *, days: int) -> pd.DataFrame:71        import yfinance as yf72 73        # yfinance periods are coarse; pick nearest.74        period = "1y" if days <= 365 else "2y" if days <= 730 else "5y"75        try:76            df = yf.download(77                sym, period=period, auto_adjust=True, threads=False, progress=False78            )79        except Exception:80            return pd.DataFrame()81        if not isinstance(df, pd.DataFrame) or df.empty:82            return pd.DataFrame()83        df = df.dropna(how="all")84        if df.empty or "Close" not in df:85            return pd.DataFrame()86        df.index = pd.to_datetime(df.index)87        return df88 89 90class EODHD(Provider):91    def __init__(self, api_token: str):92        super().__init__(name="eodhd")93        self.api_token = api_token94        self._warned = 095 96    def history(self, sym: str, *, days: int) -> pd.DataFrame:97        # EODHD uses EXCHANGE suffix, e.g. RELIANCE.NSE.98        ticker = sym.replace(".NS", ".NSE")99        to_dt = datetime.now(timezone.utc).date()100        from_dt = to_dt - timedelta(days=int(days) + 7)  # pad for weekends/holidays101        url = f"https://eodhd.com/api/eod/{ticker}"102        params = {103            "api_token": self.api_token,104            "fmt": "json",105            "period": "d",106            "from": from_dt.strftime("%Y-%m-%d"),107            "to": to_dt.strftime("%Y-%m-%d"),108        }109        try:110            r = requests.get(url, params=params, timeout=30)111            if r.status_code != 200:112                if self._warned < 3:113                    self._warned += 1114                    print(f"EODHD {ticker} -> {r.status_code}: {(r.text or '')[:200]}")115                return pd.DataFrame()116            data = r.json() or []117        except Exception:118            return pd.DataFrame()119        if not isinstance(data, list) or not data:120            if self._warned < 3:121                self._warned += 1122                s = data if isinstance(data, dict) else {"response": str(type(data))}123                print(f"EODHD {ticker} -> empty/non-list: {str(s)[:200]}")124            return pd.DataFrame()125        df = pd.DataFrame(data)126        if "date" not in df or "close" not in df:127            return pd.DataFrame()128        df["Date"] = pd.to_datetime(df["date"])129        df = df.set_index("Date").sort_index()130        out = pd.DataFrame({"Close": pd.to_numeric(df["close"], errors="coerce")})131        if "volume" in df:132            out["Volume"] = pd.to_numeric(df["volume"], errors="coerce")133        return out.dropna()134 135 136class Groww(Provider):137    def __init__(self, api_key: str, api_secret: str):138        super().__init__(name="groww")139        self._auth = GrowwAuth(api_key=api_key, api_secret=api_secret)140        self._token: str | None = None141        self._warned = 0142 143    def _token_get(self) -> str:144        if self._token:145            return self._token146        self._token = get_access_token(self._auth)147        return self._token148 149    def history(self, sym: str, *, days: int) -> pd.DataFrame:150        trading_symbol = sym.replace(".NS", "").upper()151        end = datetime.now()152        start = end - timedelta(days=int(days) + 7)  # pad for weekends/holidays153        start_s = start.strftime("%Y-%m-%d 09:15:00")154        end_s = end.strftime("%Y-%m-%d 15:30:00")155 156        try:157            candles = get_candles_range(158                self._token_get(),159                trading_symbol=trading_symbol,160                start_time=start_s,161                end_time=end_s,162                interval_in_minutes=1440,163            )164        except Exception as e:165            if self._warned < 3:166                self._warned += 1167                print(f"Groww {trading_symbol} -> {str(e)[:200]}")168            return pd.DataFrame()169 170        if not candles:171            return pd.DataFrame()172 173        df = pd.DataFrame(174            candles, columns=["ts", "Open", "High", "Low", "Close", "Volume"]175        )176        df["ts"] = (177            pd.to_datetime(df["ts"], unit="s", utc=True)178            .dt.tz_convert("Asia/Kolkata")179            .dt.tz_localize(None)180        )181        df = df.set_index("ts").sort_index()182        df.index = pd.to_datetime(df.index.date)  # normalize to date183        df = df.apply(pd.to_numeric, errors="coerce").dropna(subset=["Close"])184        return df185 186 187def _load_upstox_symbol_map() -> dict[str, str]:188    raw = (os.getenv("UPSTOX_SYMBOL_MAP") or "").strip()189    if raw:190        try:191            data = json.loads(raw)192            if isinstance(data, dict):193                return {194                    str(k).upper(): str(v) for k, v in data.items() if str(v).strip()195                }196        except Exception:197            return {}198 199    path = (os.getenv("UPSTOX_SYMBOL_MAP_FILE") or "").strip()200    if path:201        try:202            p = Path(path)203            data = json.loads(p.read_text(encoding="utf-8"))204            if isinstance(data, dict):205                return {206                    str(k).upper(): str(v) for k, v in data.items() if str(v).strip()207                }208        except Exception:209            return {}210    return {}211 212 213class Upstox(Provider):214    def __init__(self, access_token: str, symbol_map: dict[str, str]):215        super().__init__(name="upstox")216        self._token = access_token217        self._symbol_map = {k.upper(): v for k, v in symbol_map.items()}218        self._warned = 0219 220    def history(self, sym: str, *, days: int) -> pd.DataFrame:221        instrument_key = self._symbol_map.get(sym.upper())222        if not instrument_key:223            return pd.DataFrame()224 225        to_dt = datetime.now(timezone.utc).date()226        from_dt = to_dt - timedelta(days=int(days) + 7)227        encoded_key = quote(instrument_key, safe="")228        url = (229            f"https://api.upstox.com/v2/historical-candle/{encoded_key}/day/"230            f"{to_dt.strftime('%Y-%m-%d')}/{from_dt.strftime('%Y-%m-%d')}"231        )232        try:233            r = requests.get(234                url,235                headers={236                    "Accept": "application/json",237                    "Authorization": f"Bearer {self._token}",238                },239                timeout=25,240            )241            if r.status_code != 200:242                if self._warned < 3:243                    self._warned += 1244                    print(f"Upstox {sym} -> {r.status_code}: {(r.text or '')[:200]}")245                return pd.DataFrame()246            data = r.json() or {}247        except Exception:248            return pd.DataFrame()249 250        payload = data.get("data") if isinstance(data, dict) else None251        candles = payload.get("candles") if isinstance(payload, dict) else None252        if not isinstance(candles, list) or not candles:253            return pd.DataFrame()254 255        rows: list[dict[str, Any]] = []256        for c in candles:257            if not isinstance(c, list) or len(c) < 6:258                continue259            rows.append(260                {261                    "Date": pd.to_datetime(c[0]),262                    "Open": c[1],263                    "High": c[2],264                    "Low": c[3],265                    "Close": c[4],266                    "Volume": c[5],267                }268            )269 270        if not rows:271            return pd.DataFrame()272 273        df = pd.DataFrame(rows).set_index("Date").sort_index()274        df.index = pd.to_datetime(df.index).tz_localize(None)275        df.index = pd.to_datetime(df.index.date)276        return df.apply(pd.to_numeric, errors="coerce").dropna(subset=["Close"])277 278 279def _read_universe(path: Path) -> list[str]:280    out: list[str] = []281    for line in path.read_text(encoding="utf-8").splitlines():282        s = line.strip()283        if not s or s.startswith("#"):284            continue285        out.append(s.upper())286    return out287 288 289def _parquet_write(df: pd.DataFrame, out_path: Path) -> None:290    out_path.parent.mkdir(parents=True, exist_ok=True)291    con = duckdb.connect(database=":memory:")292    con.register("df", df)293    path = str(out_path).replace("'", "''")294    con.execute(f"COPY df TO '{path}' (FORMAT PARQUET)")295    con.close()296 297 298def _safe_symbol(sym: str) -> str:299    return re.sub(r"[^A-Za-z0-9._-]+", "_", sym.upper())300 301 302def _ohlcv_path(sym: str) -> Path:303    return data_root() / "ohlcv" / f"{_safe_symbol(sym)}.parquet"304 305 306def _write_ohlcv(sym: str, h: pd.DataFrame) -> None:307    if h is None or h.empty:308        return309    if "Close" not in h:310        return311    out = _ohlcv_path(sym)312    out.parent.mkdir(parents=True, exist_ok=True)313 314    df = h.copy()315    df.index = pd.to_datetime(df.index)316    df = df.reset_index()317    if "Date" not in df.columns:318        # Handle non-default index names (e.g., Groww uses "ts").319        df = df.rename(columns={df.columns[0]: "Date"})320    if "Date" not in df.columns:321        return322    for c in ["Open", "High", "Low", "Close"]:323        if c not in df.columns:324            return325    if "Volume" not in df.columns:326        df["Volume"] = 0327    df = df[["Date", "Open", "High", "Low", "Close", "Volume"]]328 329    con = duckdb.connect(database=":memory:")330    con.register("df", df)331    path = str(out).replace("'", "''")332    con.execute(f"COPY df TO '{path}' (FORMAT PARQUET)")333    con.close()334 335 336def _build_feature_table(337    universe: Iterable[str], provider: Provider, *, horizon: int, days: int338) -> pd.DataFrame:339    rows: list[pd.DataFrame] = []340    for sym in universe:341        h = provider.history(sym, days=days)342        if h.empty:343            continue344        _write_ohlcv(sym, h)345        f = build_features(h)346        if f.empty:347            continue348        y = build_targets(h["Close"], horizon=horizon).reindex(f.index)349        # Keep the latest feature rows even though their forward target is NaN.350        f = f.assign(symbol=sym, target=y)351        rows.append(f.reset_index().rename(columns={"index": "Date"}))352    return pd.concat(rows, ignore_index=True) if rows else pd.DataFrame()353 354 355def _train_ridge(df: pd.DataFrame, *, horizon: int, alpha: float) -> tuple[dict, dict]:356    df = df.copy()357    df["Date"] = pd.to_datetime(df["Date"])358    df = df.sort_values("Date").dropna()359 360    features = [c for c in df.columns if c not in {"Date", "symbol", "target"}]361 362    dates = df["Date"].drop_duplicates().sort_values()363    if len(dates) < 50:364        raise SystemExit("Not enough dates to train (need at least ~50).")365    cutoff = dates.iloc[int(len(dates) * 0.8)]366    train = df[df["Date"] <= cutoff]367    test = df[df["Date"] > cutoff]368 369    Xtr = train[features].to_numpy(dtype=float)370    ytr = train["target"].to_numpy(dtype=float)371    mu, sig = zscore_fit(Xtr)372    w = fit_ridge(zscore_apply(Xtr, mu, sig), ytr, alpha=alpha)373 374    Xte = test[features].to_numpy(dtype=float)375    yte = test["target"].to_numpy(dtype=float)376    yhat = predict(zscore_apply(Xte, mu, sig), w)377 378    ic = float(np.corrcoef(yhat, yte)[0, 1]) if len(yte) > 10 else 0.0379    hit = float((np.sign(yhat) == np.sign(yte)).mean()) if len(yte) else 0.0380 381    meta = {382        "model": "ridge",383        "horizon": horizon,384        "alpha": alpha,385        "features": features,386        "cutoff": cutoff.strftime("%Y-%m-%d"),387        "rows_train": int(len(train)),388        "rows_test": int(len(test)),389        "ic": ic,390        "hit_rate": hit,391    }392    model = {"w": w, "mu": mu, "sig": sig, "features": features}393    return meta, model394 395 396def _write_model(meta: dict, model: dict, *, model_id: str) -> tuple[str, Path]:397    v = version_id()398    out_dir = model_version_dir(model_id, v)399    out_dir.mkdir(parents=True, exist_ok=True)400    np.savez(401        out_dir / "model.npz",402        w=model["w"],403        mu=model["mu"],404        sig=model["sig"],405        features=np.array(model["features"], dtype=object),406    )407    save_meta(out_dir, meta)408    write_latest(model_id, v)409    return v, out_dir410 411 412def main() -> int:413    ap = argparse.ArgumentParser()414    ap.add_argument("--universe-file", default="data/universe/nifty50.txt")415    ap.add_argument("--universe-name", default="nifty50")416    ap.add_argument("--horizon", type=int, default=5)417    ap.add_argument("--alpha", type=float, default=10.0)418    ap.add_argument(419        "--history-days",420        type=int,421        default=365,422        help="History window to fetch per symbol",423    )424    ap.add_argument(425        "--provider",426        choices=["auto", "upstox", "groww", "eodhd", "yfinance"],427        default="auto",428    )429    ap.add_argument(430        "--upload", action="store_true", help="Upload artifacts to Supabase Storage"431    )432    ap.add_argument(433        "--prefix", default="", help="Remote prefix (default: eod/<universe-name>)"434    )435    args = ap.parse_args()436 437    universe = _read_universe(Path(args.universe_file))438    if not universe:439        raise SystemExit("Empty universe file")440 441    upstox_token = (os.getenv("UPSTOX_ACCESS_TOKEN") or "").strip()442    upstox_map = _load_upstox_symbol_map()443    eodhd_key = (os.getenv("EODHD_API_KEY") or "").strip()444    groww_key = (os.getenv("GROWW_API_KEY") or "").strip()445    groww_secret = (os.getenv("GROWW_API_SECRET") or "").strip()446    providers: list[Provider] = []447    if args.provider in ("auto", "upstox") and upstox_token and upstox_map:448        providers.append(Upstox(upstox_token, upstox_map))449    if args.provider in ("auto", "groww") and groww_key and groww_secret:450        providers.append(Groww(groww_key, groww_secret))451    if args.provider in ("auto", "eodhd") and eodhd_key:452        providers.append(EODHD(eodhd_key))453    if args.provider in ("auto", "yfinance") or not providers:454        providers.append(YFinance())455 456    if args.provider == "upstox" and (not upstox_token or not upstox_map):457        raise SystemExit("Missing UPSTOX_ACCESS_TOKEN and/or UPSTOX_SYMBOL_MAP")458    if args.provider in ("auto", "groww") and not (groww_key and groww_secret):459        print("Note: GROWW_API_KEY/GROWW_API_SECRET not set; skipping Groww.")460    if args.provider in ("auto", "eodhd") and not eodhd_key:461        print("Note: EODHD_API_KEY not set; skipping EODHD.")462    if args.provider in ("auto", "upstox") and not (upstox_token and upstox_map):463        print("Note: UPSTOX_ACCESS_TOKEN/UPSTOX_SYMBOL_MAP not set; skipping Upstox.")464 465    df = pd.DataFrame()466    used = None467    for p in providers:468        df = _build_feature_table(469            universe, p, horizon=args.horizon, days=args.history_days470        )471        if not df.empty:472            used = p.name473            break474 475    if df.empty:476        raise SystemExit("No features produced (data source unavailable).")477 478    root = data_root()479    df["Date"] = pd.to_datetime(df["Date"])480    as_of = pd.to_datetime(df["Date"].max()).strftime("%Y-%m-%d")481 482    feat_path = root / "features.parquet"483    _parquet_write(df, feat_path)484 485    model_id = f"ridge_h{args.horizon}"486    meta, model = _train_ridge(df, horizon=args.horizon, alpha=args.alpha)487    meta.update(488        {489            "universe": args.universe_name,490            "universe_size": len(universe),491            "provider": used or "",492        }493    )494    v, model_dir = _write_model(meta, model, model_id=model_id)495 496    latest = {497        "as_of_date": as_of,498        "universe": args.universe_name,499        "universe_size": len(universe),500        "horizon": args.horizon,501        "model_id": model_id,502        "model_version": v,503        "provider": used or "",504        "generated_at_utc": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"),505    }506    (root / "eod_latest.json").write_text(507        json.dumps(latest, indent=2), encoding="utf-8"508    )509 510    print(f"Wrote features -> {feat_path}")511    print(f"Wrote model -> {model_dir}")512    print(513        f"As-of {as_of} universe={args.universe_name} n={len(universe)} provider={used} model={model_id}@{v}"514    )515 516    if not args.upload:517        return 0518 519    sb = from_env(require_write=True)520    if not sb:521        raise SystemExit(522            "Missing SUPABASE_URL/SUPABASE_BUCKET/SUPABASE_SERVICE_ROLE_KEY for upload"523        )524 525    prefix = args.prefix.strip().strip("/") or f"eod/{args.universe_name}"526 527    # Upload features + model files first.528    sb.upload_bytes(529        f"{prefix}/features.parquet",530        feat_path.read_bytes(),531        content_type="application/octet-stream",532    )533    sb.upload_bytes(534        f"{prefix}/models/{model_id}/{v}/model.npz",535        (model_dir / "model.npz").read_bytes(),536        content_type="application/octet-stream",537    )538    sb.upload_bytes(539        f"{prefix}/models/{model_id}/{v}/meta.json",540        (model_dir / "meta.json").read_bytes(),541        content_type="application/json",542    )543    sb.upload_bytes(544        f"{prefix}/models/{model_id}/LATEST",545        (root / "models" / model_id / "LATEST").read_bytes(),546        content_type="text/plain",547    )548 549    # Upload per-symbol OHLCV snapshots (used by Streamlit Cloud to avoid yfinance at runtime).550    for sym in universe:551        p = _ohlcv_path(sym)552        if p.exists():553            sb.upload_bytes(554                f"{prefix}/ohlcv/{p.name}",555                p.read_bytes(),556                content_type="application/octet-stream",557            )558 559    # Publish latest.json last (\"last good snapshot\" rule).560    sb.upload_bytes(561        f"{prefix}/latest.json",562        json.dumps(latest, indent=2).encode("utf-8"),563        content_type="application/json",564    )565 566    print(f"Uploaded -> {sb.public_url(f'{prefix}/latest.json')}")567    return 0568 569 570if __name__ == "__main__":571    raise SystemExit(main())572