"""Self-contained JKP CTF submission: symmetric Löwdin PLS5 plus MVO.

Each formation-month cross-section is median-filled, winsorised at 1/99%,
standardised, and symmetrically Löwdin-orthogonalised. A five-component PLS
model predicts next-month excess returns. Estimation starts after 36 months,
uses an expanding window containing only earlier months, and refits every 12
months, matching the example model's timing.

The PLS signal replaces LightGBM in the example portfolio operator: a trailing
756-day Ledoit-Wolf covariance, ridge equal to its mean eigenvalue, and
tangency-minus-GMV weights normalised to zero net and unit gross exposure.
Stocks with fewer than 250 observed daily returns receive zero weight.

The covariance linear system is mathematically identical to the dense example
solve. Preconditioned conjugate gradients apply the covariance implicitly; the
relative tolerance is 1e-7, inside the CTF reproducibility tolerance.
"""

from __future__ import annotations

import time

import numpy as np
import pandas as pd
from scipy.sparse.linalg import LinearOperator, cg


TARGET = "ret_exc_lead1m"
DATE = "eom"
ID = "id"
TEST_FLAG = "ctff_test"

N_COMPONENTS = 5
MIN_TRAIN_MONTHS = 36
REFIT_EVERY = 12
EIGENVALUE_FLOOR = 1e-6
SAMPLE_WINDOW_DAYS = 756
MIN_DAILY_OBS = 250
MIN_NAMES = 50
DAYS_PER_MONTH = 21
KAPPA = 1.0
CG_RTOL = 1e-7
CG_MAXITER = 500


def _feature_names(features: pd.DataFrame) -> list[str]:
    if "features" not in features.columns:
        raise ValueError("features must contain a column named 'features'")
    names = features["features"].astype(str).tolist()
    if not names or len(names) != len(set(names)):
        raise ValueError("the supplied feature list is empty or contains duplicates")
    return names


def _test_mask(flag: pd.Series) -> np.ndarray:
    if pd.api.types.is_bool_dtype(flag):
        return flag.fillna(False).to_numpy(dtype=bool)
    if pd.api.types.is_numeric_dtype(flag):
        return flag.fillna(0).ne(0).to_numpy(dtype=bool)
    return (
        flag.astype(str).str.strip().str.lower()
        .isin(["1", "1.0", "true", "t"]).to_numpy(dtype=bool)
    )


def _clean_cross_section(x: np.ndarray) -> np.ndarray:
    """Clean one formation cross-section without using return availability."""
    x = np.asarray(x, dtype=np.float64).copy()
    x[~np.isfinite(x)] = np.nan
    valid = np.any(np.isfinite(x), axis=0)
    x[:, ~valid] = 0.0
    medians = np.nanmedian(x, axis=0)
    ii, jj = np.where(np.isnan(x))
    x[ii, jj] = medians[jj]
    lower, upper = np.quantile(x, [0.01, 0.99], axis=0)
    x = np.clip(x, lower, upper)
    mean = x.mean(axis=0)
    std = x.std(axis=0)
    active = std > 1e-12
    z = (x - mean) / np.where(active, std, 1.0)
    z[:, ~active] = 0.0
    return z


def _symmetric_lowdin(z: np.ndarray) -> np.ndarray:
    """Return Z (Z'Z/N)^(-1/2) with a fixed eigenvalue floor."""
    gram = z.T @ z / len(z)
    gram = (gram + gram.T) * 0.5
    eigenvalues, eigenvectors = np.linalg.eigh(gram)
    inverse_root = np.maximum(eigenvalues, EIGENVALUE_FLOOR) ** -0.5
    transform = (eigenvectors * inverse_root) @ eigenvectors.T
    return z @ transform


def _monthly_moments(x: np.ndarray, y: np.ndarray):
    """Stock-averaged moments; historical months are subsequently equal-weighted."""
    observed = np.isfinite(y)
    if not np.any(observed):
        return None
    design = np.column_stack([np.ones(int(observed.sum())), x[observed]])
    return (design.T @ design / observed.sum(),
            design.T @ y[observed] / observed.sum())


def _fit_pls5(moment: np.ndarray, cross: np.ndarray) -> np.ndarray:
    """Univariate unscaled PLS5 from centered sufficient moments."""
    mean_x = moment[0, 1:] / moment[0, 0]
    mean_y = cross[0] / moment[0, 0]
    gram = moment[1:, 1:] / moment[0, 0] - np.outer(mean_x, mean_x)
    gram = (gram + gram.T) * 0.5
    rhs = cross[1:] / moment[0, 0] - mean_x * mean_y
    coefficient = np.zeros_like(rhs)
    residual = rhs.copy()
    direction = residual.copy()
    residual_sq = float(residual @ residual)
    for _ in range(N_COMPONENTS):
        gram_direction = gram @ direction
        denominator = float(direction @ gram_direction)
        if residual_sq < 1e-24 or denominator <= 1e-24:
            break
        step = residual_sq / denominator
        coefficient += step * direction
        residual -= step * gram_direction
        new_residual_sq = float(residual @ residual)
        direction = residual + (new_residual_sq / residual_sq) * direction
        residual_sq = new_residual_sq
    return np.r_[mean_y - mean_x @ coefficient, coefficient]


def _build_daily_matrix(daily_ret: pd.DataFrame, panel_ids: np.ndarray):
    required = {"id", "date", "ret_exc"}
    if not required.issubset(daily_ret.columns):
        raise ValueError(f"daily_ret must contain {sorted(required)}")
    raw_ids = daily_ret["id"].to_numpy()
    keep = np.isin(raw_ids, panel_ids)
    raw_ids = raw_ids[keep]
    raw_dates = daily_ret["date"].to_numpy()[keep].astype("datetime64[D]")
    raw_returns = daily_ret["ret_exc"].to_numpy()[keep].astype(np.float32)
    dates, date_index = np.unique(raw_dates, return_inverse=True)
    ids, id_index = np.unique(raw_ids, return_inverse=True)
    matrix = np.full((len(dates), len(ids)), np.nan, dtype=np.float32)
    matrix[date_index, id_index] = raw_returns
    print(f"daily matrix {len(dates):,} days x {len(ids):,} ids, "
          f"{np.isfinite(matrix).mean():.1%} dense", flush=True)
    return dates, ids, matrix


def _trailing_window(matrix: np.ndarray, end: int, columns: np.ndarray):
    start = max(0, end - SAMPLE_WINDOW_DAYS)
    x = matrix[start:end, columns]
    observed = np.isfinite(x)
    enough = observed.sum(axis=0) >= MIN_DAILY_OBS
    x = np.where(observed, x, np.float32(0.0))
    x -= x.mean(axis=0, keepdims=True)
    return x, enough


def _mvo_weights(mu: np.ndarray, daily: np.ndarray):
    """Tangency-minus-GMV under the example's shrunk covariance and ridge."""
    x = np.asarray(daily * np.float32(np.sqrt(DAYS_PER_MONTH)), dtype=np.float64)
    periods, n = x.shape
    if periods < 2 or n < 2:
        raise ValueError(f"risk window is too small: {x.shape}")
    trace_sample = float(np.square(x).sum()) / periods
    mean_eigenvalue = trace_sample / n
    time_gram = x @ x.T
    norm_sample_sq = float(np.square(time_gram).sum()) / periods**2
    d2 = max(norm_sample_sq - trace_sample**2 / n, 0.0)
    row_sq = np.square(x).sum(axis=1)
    b2 = float(np.square(row_sq).sum()) / periods**2 - norm_sample_sq / periods
    delta = min(max(b2, 0.0), d2) / d2 if d2 > 1e-300 else 1.0

    # A = beta X'X + alpha I equals Sigma_LW + kappa*tr(Sigma)/n I.
    beta = (1.0 - delta) / periods
    alpha = (delta + KAPPA) * mean_eigenvalue
    diagonal = alpha + beta * np.square(x).sum(axis=0)
    operator = LinearOperator(
        (n, n),
        matvec=lambda vector: alpha * vector + beta * (x.T @ (x @ vector)),
        dtype=np.float64,
    )
    preconditioner = LinearOperator(
        (n, n), matvec=lambda vector: vector / diagonal, dtype=np.float64
    )

    def solve(rhs: np.ndarray) -> np.ndarray:
        solution, info = cg(operator, np.asarray(rhs, dtype=np.float64),
                            M=preconditioner, rtol=CG_RTOL, atol=0.0,
                            maxiter=CG_MAXITER)
        relative_residual = np.linalg.norm(operator @ solution - rhs) / max(
            np.linalg.norm(rhs), 1e-30
        )
        if info != 0 or relative_residual > 5e-6:
            raise RuntimeError(
                f"covariance solve failed: info={info}, residual={relative_residual:.3e}"
            )
        return solution

    a = solve(mu)
    one = np.ones(n, dtype=np.float64)
    b = solve(one)
    weight = a - (one @ a) / (one @ b) * b
    gross = np.abs(weight).sum()
    if not np.isfinite(gross) or gross <= 0:
        raise RuntimeError("portfolio optimisation produced invalid weights")
    return weight / gross, delta


def main(chars: pd.DataFrame, features: pd.DataFrame,
         daily_ret: pd.DataFrame) -> pd.DataFrame:
    """Return CTF portfolio weights with columns exactly id, eom, and w."""
    np.random.seed(0)
    started = time.perf_counter()
    names = _feature_names(features)
    required = [ID, DATE, TARGET, *names]
    missing = [column for column in required if column not in chars.columns]
    if missing:
        raise ValueError(f"chars is missing required columns, e.g. {missing[:5]}")

    ids = pd.to_numeric(chars[ID], errors="raise").to_numpy(dtype=np.int64)
    dates = pd.to_datetime(chars[DATE], errors="raise").to_numpy(dtype="datetime64[ns]")
    test = (_test_mask(chars[TEST_FLAG]) if TEST_FLAG in chars.columns
            else np.zeros(len(chars), dtype=bool))
    use_test = bool(test.any())
    order = np.lexsort((ids, dates.view(np.int64)))
    ordered_dates = dates[order]
    boundaries = np.r_[0, np.flatnonzero(
        ordered_dates[1:] != ordered_dates[:-1]) + 1, len(order)]

    daily_dates, daily_ids, daily_matrix = _build_daily_matrix(
        daily_ret, np.unique(ids)
    )
    mapped = np.clip(np.searchsorted(daily_ids, ids), 0, len(daily_ids) - 1)
    has_daily = daily_ids[mapped] == ids

    p = len(names)
    moment_sum = np.zeros((p + 1, p + 1), dtype=np.float64)
    cross_sum = np.zeros(p + 1, dtype=np.float64)
    labelled_months = 0
    coefficient: np.ndarray | None = None
    outputs: list[pd.DataFrame] = []

    for month_number, (left, right) in enumerate(
            zip(boundaries[:-1], boundaries[1:]), start=1):
        idx = order[left:right]
        month = ordered_dates[left]
        raw = chars.iloc[idx][names].to_numpy(dtype=np.float64, copy=True)
        transformed = _symmetric_lowdin(_clean_cross_section(raw))
        formation_index = month_number - 1

        if formation_index >= MIN_TRAIN_MONTHS:
            if (formation_index - MIN_TRAIN_MONTHS) % REFIT_EVERY == 0:
                coefficient = _fit_pls5(
                    moment_sum / labelled_months, cross_sum / labelled_months
                )
            if coefficient is None:
                raise RuntimeError("PLS coefficient unavailable after training window")
            keep = test[idx] if use_test else np.ones(len(idx), dtype=bool)
            if np.any(keep):
                scores = coefficient[0] + transformed @ coefficient[1:]
                present = has_daily[idx]
                daily_columns = mapped[idx][present]
                end = int(np.searchsorted(
                    daily_dates, month.astype("datetime64[D]"), side="right"
                ))
                daily_window, enough = _trailing_window(
                    daily_matrix, end, daily_columns
                )
                usable = np.zeros(len(idx), dtype=bool)
                usable[np.flatnonzero(present)[enough]] = True
                if int(usable.sum()) < MIN_NAMES:
                    raise RuntimeError(
                        f"month {pd.Timestamp(month):%Y-%m}: only "
                        f"{int(usable.sum())} usable securities"
                    )
                weights = np.zeros(len(idx), dtype=np.float64)
                weights[usable], delta = _mvo_weights(
                    scores[usable], daily_window[:, enough]
                )
                outputs.append(pd.DataFrame({
                    ID: ids[idx][keep], DATE: pd.to_datetime(dates[idx][keep]),
                    "w": weights[keep],
                }))
                print(f"{pd.Timestamp(month):%Y-%m} names={len(idx):,} "
                      f"held={int(usable.sum()):,} delta={delta:.3f}", flush=True)

        y = pd.to_numeric(chars.iloc[idx][TARGET], errors="coerce").to_numpy(
            dtype=np.float64
        )
        monthly = _monthly_moments(transformed, y)
        if monthly is not None:
            moment, cross = monthly
            moment_sum += moment
            cross_sum += cross
            labelled_months += 1

    if not outputs:
        raise RuntimeError(
            f"no weights produced: need at least {MIN_TRAIN_MONTHS + 1} months"
        )
    out = pd.concat(outputs, ignore_index=True)
    out = out[[ID, DATE, "w"]].sort_values(
        [DATE, ID], kind="stable").reset_index(drop=True)
    out[ID] = out[ID].astype(np.int64)
    out["w"] = out["w"].astype(np.float64)
    if out.isna().any().any():
        raise RuntimeError("output contains missing values")
    print(f"completed {len(out):,} rows in "
          f"{(time.perf_counter()-started)/60:.1f} min", flush=True)
    return out


if __name__ == "__main__":
    # Local convenience only. The CTF imports main() and ignores this block.
    from pathlib import Path

    import pyarrow as pa
    import pyarrow.parquet as pq

    here = Path(__file__).resolve().parent
    data = here.parent / "data"
    feature_frame = pd.read_parquet(data / "ctff_features.parquet")
    feature_names = _feature_names(feature_frame)
    selected = [ID, DATE, TARGET, TEST_FLAG, *feature_names]
    table = pq.read_table(data / "ctff_chars.parquet", columns=selected)
    schema = pa.schema([
        pa.field(field.name, pa.float32()) if field.name in feature_names else field
        for field in table.schema
    ])
    chars_frame = table.cast(schema).to_pandas(self_destruct=True)
    daily_frame = pq.read_table(data / "ctff_daily_ret.parquet").to_pandas(
        date_as_object=False, self_destruct=True
    )
    result = main(chars_frame, feature_frame, daily_frame)
    output_path = here / "output_pls5_mvo.csv"
    result.to_csv(output_path, index=False, date_format="%Y-%m-%d")
    print(f"wrote {output_path} ({output_path.stat().st_size/1_048_576:.1f} MB)")
