CoolFace
Apppublic

shyameati/transcripts-api

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
app.py193 linesDownload Raw Back to root
1import logging2from fastapi import FastAPI, HTTPException3from fastapi.responses import HTMLResponse4from fastapi.middleware.cors import CORSMiddleware5from datasets import load_dataset, load_from_disk6import numpy as np7import os8 9app = FastAPI()10 11# ---------------------------------------------------------12# Logging13# ---------------------------------------------------------14logging.basicConfig(15    level=logging.INFO,16    format="%(asctime)s [%(levelname)s] %(message)s"17)18logger = logging.getLogger(__name__)19 20# ---------------------------------------------------------21# CORS22# ---------------------------------------------------------23app.add_middleware(24    CORSMiddleware,25    allow_origins=["*"],26    allow_methods=["*"],27    allow_headers=["*"],28)29 30# ---------------------------------------------------------31# Dataset caching configuration32# ---------------------------------------------------------33DATASET_NAME = "kurry/sp500_earnings_transcripts"34CACHE_PATH = "/data/hf_dataset"   # persistent bucket mount35dataset_cache = None36 37 38def load_hf_dataset():39    """40    Loads the HF dataset with persistent caching.41    - If /data/hf_dataset exists → load from disk (fast, offline)42    - Else → download once, save to disk, then load43    """44    global dataset_cache45 46    if dataset_cache is not None:47        return dataset_cache48 49    if os.path.exists(CACHE_PATH):50        logger.info(f"Loading dataset from cache at {CACHE_PATH}")51        dataset_cache = load_from_disk(CACHE_PATH)52        logger.info(f"Loaded {len(dataset_cache)} rows from cached dataset")53        return dataset_cache54 55    logger.info(f"Downloading HF dataset: {DATASET_NAME}")56    ds = load_dataset(DATASET_NAME, split="train")57 58    logger.info(f"Saving dataset to cache at {CACHE_PATH}")59    ds.save_to_disk(CACHE_PATH)60 61    dataset_cache = ds62    logger.info(f"Dataset cached and loaded ({len(ds)} rows)")63    return ds64 65 66# ---------------------------------------------------------67# JSON-safe conversion68# ---------------------------------------------------------69def to_json_safe(obj):70    if isinstance(obj, (np.integer,)):71        return int(obj)72    if isinstance(obj, (np.floating,)):73        return float(obj)74    if isinstance(obj, (np.ndarray, list)):75        return [to_json_safe(x) for x in obj]76    if isinstance(obj, dict):77        return {k: to_json_safe(v) for k, v in obj.items()}78    return obj79 80 81# ---------------------------------------------------------82# Serve index.html83# ---------------------------------------------------------84@app.get("/", response_class=HTMLResponse)85def serve_index():86    if not os.path.exists("index.html"):87        return "<h1>index.html not found</h1>"88    with open("index.html", "r") as f:89        return f.read()90 91 92# ---------------------------------------------------------93# List all symbols94# ---------------------------------------------------------95@app.get("/tickers")96def get_tickers():97    ds = load_hf_dataset()98    symbols = sorted(set([s.upper() for s in ds["symbol"]]))99    return {"tickers": symbols}100 101 102# ---------------------------------------------------------103# Get transcript for a symbol104# ---------------------------------------------------------105@app.get("/transcript/{symbol}")106def get_transcript(symbol: str):107    ds = load_hf_dataset()108    symbol = symbol.upper()109 110    logger.info(f"Fetching transcript for: {symbol}")111 112    rows = [r for r in ds if r["symbol"].upper() == symbol]113 114    if not rows:115        raise HTTPException(status_code=404, detail=f"No transcript found for {symbol}")116 117    safe_rows = [to_json_safe(r) for r in rows]118 119    return {"symbol": symbol, "records": safe_rows}120 121 122# ---------------------------------------------------------123# Dataset info (size + columns)124# ---------------------------------------------------------125@app.get("/dataset-info")126def dataset_info():127    ds = load_hf_dataset()128 129    info = {130        "num_rows": len(ds),131        "columns": ds.column_names,132        "cache_path": CACHE_PATH,133    }134 135    return info136 137 138# ---------------------------------------------------------139# Dataset summary (high-level stats)140# ---------------------------------------------------------141@app.get("/dataset-summary")142def dataset_summary():143    ds = load_hf_dataset()144 145    symbols = set([s.upper() for s in ds["symbol"]])146    years = set(ds["year"])147    quarters = set(ds["quarter"])148 149    dates = [d for d in ds["date"] if d is not None]150    min_date = min(dates) if dates else None151    max_date = max(dates) if dates else None152 153    summary = {154        "total_rows": len(ds),155        "unique_symbols": len(symbols),156        "symbols_sample": sorted(list(symbols))[:20],157        "year_range": {158            "min_year": min(years),159            "max_year": max(years)160        },161        "quarters_present": sorted(list(quarters)),162        "date_range": {163            "min_date": min_date,164            "max_date": max_date165        },166        "company_count": len(set(ds["company_id"])),167    }168 169    return summary170 171@app.get("/check/{symbol}")172def check_symbol(symbol: str):173    ds = load_hf_dataset()174    symbol = symbol.upper()175 176    exists = any(r["symbol"].upper() == symbol for r in ds)177 178    if not exists:179        logger.warning(f"Symbol not found: {symbol}")180        return {181            "symbol": symbol,182            "exists": False,183            "message": f"Symbol '{symbol}' does not exist in the dataset."184        }185 186    logger.info(f"Symbol exists: {symbol}")187    return {188        "symbol": symbol,189        "exists": True,190        "message": f"Symbol '{symbol}' exists in the dataset."191    }192 193