"""
Cross-sectional dollar-neutral long–short (Medallion-style *direction*, not their IP).

Uses the same row alignment as `daily_sequence_*`: meta row i matches y[i].

Long top N names by signal, short bottom N by signal, equal dollar each leg:
  r_port = mean(r_long) - mean(r_short)

Weights sum to 0; gross exposure is 2× notional if |long| = |short| = N.
"""

from __future__ import annotations

import os
from typing import Literal

import numpy as np
import pandas as pd

Y_FILE = {
    1: "y_cc",
    3: "y_cc_3d",
    5: "y_cc_5d",
}


def load_forward_return_panel(prefix: str, hold_days: Literal[1, 3, 5] = 1) -> pd.DataFrame:
    """symbol, asof_date, fwd_ret aligned to daily_sequence rows."""
    key = Y_FILE.get(hold_days)
    if key is None:
        raise ValueError(f"hold_days must be 1, 3, or 5; got {hold_days}")
    meta_path = f"{prefix}_meta.parquet"
    y_path = f"{prefix}_{key}.npy"
    if not os.path.isfile(meta_path):
        raise FileNotFoundError(meta_path)
    if not os.path.isfile(y_path):
        raise FileNotFoundError(y_path)
    meta = pd.read_parquet(meta_path)
    y = np.load(y_path, mmap_mode="r")
    if len(meta) != len(y):
        raise ValueError(f"meta rows {len(meta)} != y len {len(y)}")
    out = pd.DataFrame(
        {
            "symbol": meta["symbol"].astype(str),
            "asof_date": pd.to_datetime(meta["asof_date"]),
            "fwd_ret": np.asarray(y, dtype=np.float64),
        }
    )
    return out


def attach_signal(
    panel: pd.DataFrame,
    *,
    source: Literal["parquet", "random"],
    predictions_path: str | None = None,
    signal_col: str = "pred_cc_reg",
    random_seed: int = 42,
) -> pd.DataFrame:
    """Add column `signal` to panel (copy)."""
    df = panel.copy()
    if source == "parquet":
        if not predictions_path or not os.path.isfile(predictions_path):
            raise FileNotFoundError(f"predictions parquet: {predictions_path}")
        pred = pd.read_parquet(predictions_path)
        if signal_col not in pred.columns:
            raise KeyError(f"{signal_col} not in {predictions_path}; columns: {list(pred.columns)}")
        pred = pred.copy()
        pred["asof_date"] = pd.to_datetime(pred["asof_date"])
        pred["symbol"] = pred["symbol"].astype(str)
        sub = pred[["symbol", "asof_date", signal_col]].rename(columns={signal_col: "signal"})
        df = df.merge(sub, on=["symbol", "asof_date"], how="inner")
    elif source == "random":
        rng = np.random.default_rng(int(random_seed))
        df["signal"] = rng.standard_normal(len(df))
    else:
        raise ValueError(source)
    return df


def cross_sectional_long_short_daily(
    df: pd.DataFrame,
    *,
    signal_col: str = "signal",
    ret_col: str = "fwd_ret",
    top_n: int = 25,
    bottom_n: int = 25,
    min_names: int = 80,
    date_col: str = "asof_date",
) -> pd.DataFrame:
    """
    Per calendar date: long top_n by signal, short bottom_n by signal.
    Skip days with fewer than min_names valid rows.
    """
    rows: list[dict] = []
    for dt, g in df.groupby(date_col, sort=True):
        g = g.dropna(subset=[signal_col, ret_col])
        if len(g) < max(min_names, top_n + bottom_n + 2):
            continue
        g = g.sort_values(signal_col, ascending=False, kind="mergesort")
        longs = g.head(int(top_n))
        shorts = g.tail(int(bottom_n))
        r_l = float(longs[ret_col].mean())
        r_s = float(shorts[ret_col].mean())
        port = r_l - r_s
        rows.append(
            {
                "asof_date": dt,
                "n_eligible": int(len(g)),
                "n_long": int(len(longs)),
                "n_short": int(len(shorts)),
                "mean_ret_long": r_l,
                "mean_ret_short": r_s,
                "port_ret_gross": port,
            }
        )
    return pd.DataFrame(rows)


def apply_round_trip_cost(daily: pd.Series, bps: float) -> pd.Series:
    cost = 2.0 * float(bps) / 10_000.0
    return daily.astype(np.float64) - cost


def summarize_daily_returns(port_ret: pd.Series) -> dict[str, float]:
    r = port_ret.astype(np.float64)
    if len(r) < 2:
        return {"n_days": float(len(r)), "mean_daily": float(r.mean()), "sharpe_252": float("nan")}
    mu = float(r.mean())
    sd = float(r.std(ddof=1))
    sharpe = (mu / sd * np.sqrt(252)) if sd > 1e-12 else float("nan")
    cum = float(np.prod(1.0 + r) - 1.0)
    return {
        "n_days": float(len(r)),
        "mean_daily": mu,
        "std_daily": sd,
        "sharpe_252": sharpe,
        "cum_return": cum,
    }
