from __future__ import annotations

import gc
import os
import sys
from dataclasses import dataclass
from datetime import date, time
import pandas as pd

PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
sys.path.append(PROJECT_ROOT)

OFFSET = 30
TRAIN_END_DATE = date(2026, 4, 1)
DATA_DIR = os.path.join(PROJECT_ROOT, "2026_bk_data")


@dataclass
class BKStockData:
    X_train: pd.DataFrame
    X_test: pd.DataFrame
    y_train: pd.Series
    y_test: pd.Series
    meta_train: pd.DataFrame
    meta_test: pd.DataFrame

def calculate_forward_return(group: pd.DataFrame) -> pd.DataFrame:
    """Calculate same-day forward return with a fixed minute offset."""
    group = group.sort_values("timestamps").copy()
    group["return_30min"] = group["close"].shift(-OFFSET) / group["close"] - 1
    return group


def build_stock_frame(code: str, stock_df: pd.DataFrame) -> pd.DataFrame:
    stock_df["timestamps"] = pd.to_datetime(stock_df["timestamps"])
    stock_df = stock_df.sort_values("timestamps").reset_index(drop=True)

    cutoff = time(14, 60 - OFFSET)
    feature_cols = [
        col for col in stock_df.columns if col not in {"timestamps", "close", "y_1return"}
    ]

    X = (
        stock_df.loc[stock_df["timestamps"].dt.time <= cutoff, ["timestamps", *feature_cols]]
        .copy()
        .reset_index(drop=True)
    )
    X[feature_cols] = X[feature_cols].astype("float32")

    close = stock_df.loc[:, ["timestamps", "close"]].copy()
    close["date"] = close["timestamps"].dt.date
    close = close.groupby("date", group_keys=False).apply(
        calculate_forward_return,
        include_groups=False,
    )
    close = close.loc[close["timestamps"].dt.time <= cutoff, ["timestamps", "return_30min"]]
    close = close.reset_index(drop=True)

    merged = X.merge(close, on="timestamps", how="inner")
    merged.insert(0, "code", code)
    merged["target"] = merged["return_30min"] * 100
    return merged


def load_bk_stock_dataset(
    industry_code: str,
) -> BKStockData:
    data_path = os.path.join(DATA_DIR, f"{industry_code}.pkl")

    data_by_stock = pd.read_pickle(data_path)
    codes = sorted(data_by_stock)

    merged_frames = [build_stock_frame(code, data_by_stock[code]) for code in codes]
    del data_by_stock
    gc.collect()

    data = pd.concat([df for df in merged_frames if not df.empty], ignore_index=True)
    del merged_frames
    gc.collect()

    data = data.sort_values("timestamps").reset_index(drop=True)
    meta = data[["code", "timestamps", "target"]].reset_index(drop=True)

    feature_cols = [
        col
        for col in data.columns
        if col not in {"code", "timestamps", "close", "y_1return", "return_30min", "target"}
    ]
    X = data[feature_cols]
    y = data["target"]
    del data
    gc.collect()

    train_mask = meta["timestamps"].dt.date < TRAIN_END_DATE
    test_mask = ~train_mask

    if not train_mask.any():
        raise ValueError(f"no train samples found before {TRAIN_END_DATE}")
    if not test_mask.any():
        raise ValueError(f"no test samples found on or after {TRAIN_END_DATE}")

    X_train = X.loc[train_mask].copy().reset_index(drop=True)
    X_test = X.loc[test_mask].copy().reset_index(drop=True)
    y_train = y.loc[train_mask].copy().reset_index(drop=True)
    y_test = y.loc[test_mask].copy().reset_index(drop=True)
    meta_train = meta.loc[train_mask].copy().reset_index(drop=True)
    meta_test = meta.loc[test_mask].copy().reset_index(drop=True)
    del X, y
    gc.collect()

    test_time_mask = meta_test["timestamps"].dt.minute.isin([0, 30])
    X_test = X_test.loc[test_time_mask].reset_index(drop=True)
    y_test = y_test.loc[test_time_mask].reset_index(drop=True)
    meta_test = meta_test.loc[test_time_mask].reset_index(drop=True)


    return BKStockData(
        X_train=X_train,
        X_test=X_test,
        y_train=y_train,
        y_test=y_test,
        meta_train=meta_train,
        meta_test=meta_test,
    )
