shyameati/transcripts-api
0
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 