import os
import torch
import numpy as np
import pandas as pd
from torch.utils.data import IterableDataset
from dolphindb_tools.dataloader import DDBDataLoader
import dolphindb as ddb


class DolphinDBStreamingDataset(IterableDataset):
    """
    基于 DolphinDB 原生 DDBDataLoader 的流式数据集。
    直接产出 (batch_size, window_len, num_factors) 的训练样本。
    """

    def __init__(self, data_type: str, config: dict, rank: int = None, world_size: int = None):
        """
        data_type: 'train' 或 'val'
        config: 包含 ddb_host, ddb_port, ddb_user, ddb_password, ddb_table,
                train_time_range, val_time_range, lookback_window, predict_window,
                feature_list, time_feature_list, batch_size, seed 等。
        """
        self.config = config
        self.data_type = data_type
        self.window = config['lookback_window'] + config['predict_window'] + 1
        self.time_feature_list = config['time_feature_list']

        # 获取分布式参数（兼容 torchrun 环境变量）
        if rank is None:
            rank = int(os.environ.get('LOCAL_RANK', 0))
        if world_size is None:
            world_size = int(os.environ.get('WORLD_SIZE', 1))
        self.rank = rank
        self.world_size = world_size

        # 1. 创建 DolphinDB Session（注意：DDBDataLoader 会使用该 session）
        self.sess = ddb.Session()
        self.sess.connect(config['ddb_host'], config['ddb_port'],
                          config['ddb_user'], config['ddb_password'])

        # 2. 时间范围
        if data_type == 'train':
            start, end = config['train_time_range']
        else:
            start, end = config['val_time_range']
        self.start_str = pd.Timestamp(start).strftime('%Y.%m.%d')
        self.end_str   = pd.Timestamp(end).strftime('%Y.%m.%d')

        # 3. 获取该时间段内的所有股票代码
        stock_sql = f"""
            exec distinct code
            from {config['ddb_table']}
            where trade_date = 2026.04.20
        """
        all_stocks = self.sess.run(stock_sql).tolist()[:100]
        print(f"[Rank {rank}] Total stocks in {data_type} set: {len(all_stocks)}")

        # 4. 按 rank 均匀分片（丢弃余数，保证每个进程股票数量相同）
        per_rank = len(all_stocks) // world_size
        if per_rank == 0:
            raise RuntimeError(f"world_size ({world_size}) > number of stocks ({len(all_stocks)}).")
        start_idx = rank * per_rank
        end_idx = start_idx + per_rank
        rank_stocks = all_stocks[start_idx:end_idx]
        self.stock_list = [f"'{code}'" for code in rank_stocks]
        print(f"[Rank {rank}] Assigned {len(self.stock_list)} stocks (indices {start_idx}-{end_idx-1})")

        # 4. 获取月份（用于分区取数,如果不设置额外分区可忽略）
        month_sql = f"""
            exec distinct month(trade_date)
            from {config['ddb_table']}
            where trade_date >= {self.start_str}
                and trade_date < {self.end_str}
            """
        self.month_cols = self.sess.run(month_sql).tolist()
        print(f"自动获取 {len(self.month_cols)} 个月份")

        # 5. 从ddb取数sql，在取数过程中直接处理时间特征，减少后续处理开销
        self.sql = f"""select code, trade_time.minuteOfHour() as minute, trade_time.hour() as hour, trade_date.weekday() as weekday, trade_date.dayOfMonth() as day, trade_date.monthOfYear() as month, open, high, low, close, vol, amount from loadTable("dfs://tushare_minute_db", "stock_1min_k") where trade_date >= {self.start_str} and trade_date < {self.end_str} """

        # 6. 创建 DDBDataLoader
        self.ddb_loader = DDBDataLoader(
            ddbSession=self.sess,
            sql=self.sql,
            targetCol=self.time_feature_list, 
            inputCol=config['feature_list'],
            batchSize=config['batch_size'],
            windowSize=[self.window,self.window],
            windowStride=[10,10],
            groupCol="code",
            groupScheme=self.stock_list,
            seed=config['seed'] + rank,
            dropLast=False,
            device='cuda',
            prefetchBatch=2,
            groupPoolSize=3
        )

    def __iter__(self):
        """
        迭代返回 (x_tensor, _)，其中 x_tensor 形状为
        (batch_size, window_len, num_factors)
        """
        for batch_x, batch_y in self.ddb_loader:
            yield batch_x, batch_y 