#!/usr/bin/env python3
"""
OmmAlpha fundamental-data fetcher.

Input CSV columns: stock_id,symbol,exchange
Output JSON: current fundamentals + reported quarterly EPS history.

Provider: Yahoo Finance via yfinance. This is intended as a practical,
no-API-key research feed. OmmAlpha stores the fetched snapshot so stock pages
never make external requests during normal browsing.
"""
from __future__ import annotations

import argparse
import json
import math
import sys
from datetime import date, datetime, timezone
from pathlib import Path
from typing import Any

import pandas as pd

try:
    import yfinance as yf
except Exception as exc:  # pragma: no cover - server dependency check
    print(json.dumps({"status": "failed", "error": f"yfinance import failed: {exc}"}), file=sys.stderr)
    raise


def args() -> argparse.Namespace:
    p = argparse.ArgumentParser()
    p.add_argument("--input", required=True)
    p.add_argument("--out", required=True)
    return p.parse_args()


def finite_number(value: Any) -> float | None:
    if value is None:
        return None
    try:
        v = float(value)
    except Exception:
        return None
    return v if math.isfinite(v) else None


def text_value(value: Any) -> str | None:
    if value is None:
        return None
    s = str(value).strip()
    return s if s else None


def pct_from_fraction(value: Any) -> float | None:
    v = finite_number(value)
    return None if v is None else v * 100.0


def provider_symbol(symbol: str, exchange: str) -> str:
    symbol = symbol.strip().upper()
    exchange = exchange.strip().upper()
    if symbol.endswith((".NS", ".BO")):
        return symbol
    if exchange == "BSE":
        return f"{symbol}.BO"
    return f"{symbol}.NS"


def dataframe_row(df: pd.DataFrame | None, candidates: list[str]) -> pd.Series | None:
    if df is None or df.empty:
        return None
    normalized = {str(idx).lower().replace(" ", "").replace("_", ""): idx for idx in df.index}
    for name in candidates:
        key = name.lower().replace(" ", "").replace("_", "")
        if key in normalized:
            return df.loc[normalized[key]]
    return None


def eps_from_income_stmt(ticker: Any) -> list[dict[str, Any]]:
    rows: list[dict[str, Any]] = []
    try:
        stmt = ticker.get_income_stmt(freq="quarterly", pretty=True)
    except Exception:
        try:
            stmt = ticker.quarterly_income_stmt
        except Exception:
            return rows

    eps_series = dataframe_row(stmt, ["Diluted EPS", "Basic EPS", "DilutedEPS", "BasicEPS"])
    if eps_series is None:
        return rows

    for col, value in eps_series.items():
        eps = finite_number(value)
        if eps is None:
            continue
        try:
            d = pd.Timestamp(col).date().isoformat()
        except Exception:
            continue
        rows.append({"period_end": d, "eps_actual": eps, "source": "YF_INCOME_STATEMENT"})

    rows.sort(key=lambda x: x["period_end"])
    return rows


def eps_from_earnings_dates(ticker: Any) -> list[dict[str, Any]]:
    rows: list[dict[str, Any]] = []
    try:
        ed = ticker.get_earnings_dates(limit=24)
    except Exception:
        return rows

    if ed is None or getattr(ed, "empty", True):
        return rows

    eps_col = None
    for name in ("Reported EPS", "reportedEPS", "epsActual", "EPS Actual"):
        if name in ed.columns:
            eps_col = name
            break
    if eps_col is None:
        return rows

    for idx, r in ed.iterrows():
        eps = finite_number(r.get(eps_col))
        if eps is None:
            continue
        try:
            d = pd.Timestamp(idx).date().isoformat()
        except Exception:
            continue
        rows.append({"period_end": d, "eps_actual": eps, "source": "YF_EARNINGS_DATES"})

    # Some feeds may return multiple entries around one fiscal quarter.
    by_date: dict[str, dict[str, Any]] = {r["period_end"]: r for r in rows}
    out = list(by_date.values())
    out.sort(key=lambda x: x["period_end"])
    return out


def add_eps_growth(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
    rows = sorted(rows, key=lambda x: x["period_end"])
    for i, row in enumerate(rows):
        row["eps_yoy_growth_pct"] = None
        if i < 4:
            continue
        current = finite_number(row.get("eps_actual"))
        prior = finite_number(rows[i - 4].get("eps_actual"))
        # Avoid meaningless explosions around near-zero prior EPS.
        if current is None or prior is None or abs(prior) < 0.05:
            continue
        row["eps_yoy_growth_pct"] = ((current / prior) - 1.0) * 100.0
    return rows


def fetch_one(stock_id: int, symbol: str, exchange: str) -> dict[str, Any]:
    ps = provider_symbol(symbol, exchange)
    t = yf.Ticker(ps)

    try:
        info = t.get_info() or {}
    except Exception:
        info = {}

    eps_rows = eps_from_earnings_dates(t)
    if len(eps_rows) < 5:
        fallback = eps_from_income_stmt(t)
        by_date = {r["period_end"]: r for r in eps_rows}
        for r in fallback:
            by_date.setdefault(r["period_end"], r)
        eps_rows = sorted(by_date.values(), key=lambda x: x["period_end"])

    eps_rows = add_eps_growth(eps_rows)

    trailing_eps = finite_number(info.get("trailingEps"))
    if trailing_eps is None and len(eps_rows) >= 4:
        vals = [finite_number(r.get("eps_actual")) for r in eps_rows[-4:]]
        if all(v is not None for v in vals):
            trailing_eps = float(sum(v for v in vals if v is not None))

    current = {
        "provider_symbol": ps,
        "as_of_date": date.today().isoformat(),
        "currency": text_value(info.get("currency")) or "INR",
        "market_cap": finite_number(info.get("marketCap")),
        "enterprise_value": finite_number(info.get("enterpriseValue")),
        "trailing_pe": finite_number(info.get("trailingPE")),
        "forward_pe": finite_number(info.get("forwardPE")),
        "price_to_book": finite_number(info.get("priceToBook")),
        "eps_ttm": trailing_eps,
        "forward_eps": finite_number(info.get("forwardEps")),
        "book_value_per_share": finite_number(info.get("bookValue")),
        "revenue_ttm": finite_number(info.get("totalRevenue")),
        "net_income_ttm": finite_number(info.get("netIncomeToCommon")),
        "roe_pct": pct_from_fraction(info.get("returnOnEquity")),
        "roa_pct": pct_from_fraction(info.get("returnOnAssets")),
        # Yahoo exposes debtToEquity as percentage points (e.g. 42.7 = 42.7%).
        "debt_to_equity_pct": finite_number(info.get("debtToEquity")),
        "current_ratio": finite_number(info.get("currentRatio")),
        "operating_margin_pct": pct_from_fraction(info.get("operatingMargins")),
        "profit_margin_pct": pct_from_fraction(info.get("profitMargins")),
        "revenue_growth_yoy_pct": pct_from_fraction(info.get("revenueGrowth")),
        "earnings_growth_yoy_pct": pct_from_fraction(info.get("earningsGrowth")),
        "dividend_yield_pct": pct_from_fraction(info.get("dividendYield")),
        "payout_ratio_pct": pct_from_fraction(info.get("payoutRatio")),
        "beta": finite_number(info.get("beta")),
        "shares_outstanding": finite_number(info.get("sharesOutstanding")),
        "sector": text_value(info.get("sector")),
        "industry": text_value(info.get("industry")),
        "source": "Yahoo Finance via yfinance",
        "fetched_at_utc": datetime.now(timezone.utc).isoformat(),
    }

    useful = sum(v is not None for k, v in current.items() if k not in {"provider_symbol", "as_of_date", "currency", "source", "fetched_at_utc"})
    status = "ok" if useful > 0 or eps_rows else "empty"

    return {
        "stock_id": stock_id,
        "symbol": symbol,
        "exchange": exchange,
        "status": status,
        "current": current,
        "eps_history": eps_rows,
    }


def main() -> None:
    a = args()
    inp = Path(a.input)
    out = Path(a.out)

    df = pd.read_csv(inp)
    required = {"stock_id", "symbol", "exchange"}
    missing = required - set(df.columns)
    if missing:
        raise RuntimeError(f"Input CSV missing columns: {sorted(missing)}")

    results: list[dict[str, Any]] = []
    errors: list[dict[str, Any]] = []

    for _, r in df.iterrows():
        stock_id = int(r["stock_id"])
        symbol = str(r["symbol"]).strip().upper()
        exchange = str(r.get("exchange", "NSE")).strip().upper()
        try:
            results.append(fetch_one(stock_id, symbol, exchange))
        except Exception as exc:
            errors.append({"stock_id": stock_id, "symbol": symbol, "error": str(exc)})

    payload = {
        "status": "completed",
        "provider": "Yahoo Finance via yfinance",
        "yfinance_version": getattr(yf, "__version__", "unknown"),
        "requested": int(len(df)),
        "returned": len(results),
        "errors": errors,
        "results": results,
    }
    out.parent.mkdir(parents=True, exist_ok=True)
    out.write_text(json.dumps(payload, indent=2, allow_nan=False), encoding="utf-8")
    print(json.dumps({k: payload[k] for k in ["status", "provider", "yfinance_version", "requested", "returned"]}, indent=2))


if __name__ == "__main__":
    main()
