import dolphindb as ddb
import torch
from torch.utils.data import IterableDataset, DataLoader

class DolphinDBPartitionForDataset(IterableDataset):
    def __init__(self, host, port, username, password, db_path, table_name, partition_col):
        """
        基于 For 循环按分区动态拉取的数据集
        :param db_path: 分布式数据库路径，如 'dfs://HighFreqDB'
        :param table_name: 分布式表名，如 'tick_data'
        :param partition_col: 用于切分查询的分区列名，如 'date'
        """
        super().__init__()
        self.host = host
        self.port = port
        self.username = username
        self.password = password
        self.db_path = db_path
        self.table_name = table_name
        self.partition_col = partition_col
        
        # 1. 初始化时，先获取该表所有现存的分区值（例如所有有数据的日期）
        print("正在连接 DolphinDB 获取分区列表...")
        s = ddb.Session()
        s.connect(self.host, self.port, self.username, self.password)
        
        # 通过执行一小段脚本，把分区列的唯一值取出来（通常是几十到几百个日期，数据量极小）
        script = f"""
            t = loadTable("{self.db_path}", "{self.table_name}")
            exec distinct({self.partition_col}) from t
        """
        # 拿到分区值列表，例如: [2026.01.02, 2026.01.03, ...]
        self.partitions = list(s.run(script))
        s.close()
        
        print(f"成功！检测到该表共有 {len(self.partitions)} 个不同的分区。")

    def __iter__(self):
        # 2. 获取 PyTorch 多进程信息
        worker_info = torch.utils.data.get_worker_info()
        
        if worker_info is None:
            # 单进程模式：处理所有分区
            my_partitions = self.partitions
        else:
            # 多进程模式：利用取模算法，把分区（日期）均匀分给不同的 CPU 子进程
            my_partitions = [
                p for i, p in enumerate(self.partitions) 
                if i % worker_info.num_workers == worker_info.id
            ]

        # 3. 为当前子进程建立独立的子连接
        s = ddb.Session()
        s.connect(self.host, self.port, self.username, self.password)

        try:
            # 4. 用 FOR 循环逐个遍历当前进程负责的分区
            for part in my_partitions:
                
                # 针对 DolphinDB 不同的分区字段类型，拼接 SQL 时注意格式
                # 如果 partition_col 是 DATE 类型，DolphinDB 接受 2026.01.02 这种格式
                sql_query = f"""
                    t = loadTable("{self.db_path}", "{self.table_name}")
                    select feat1, feat2, label from t where {self.partition_col} = {part}
                """
                
                # 真正从远端拉取【这一个分区】的数据到本地内存
                df = s.run(sql_query)
                
                if df is None or len(df) == 0:
                    continue
                
                # 5. 【向量化优化】整块转成 Tensor
                # (这里的 'feat1', 'feat2', 'label' 记得根据你的实际列名修改)
                features_tensor = torch.tensor(df[['feat1', 'feat2']].values, dtype=torch.float32)
                labels_tensor = torch.tensor(df['label'].values, dtype=torch.float32)
                
                # 6. 【分区内打乱】
                perm = torch.randperm(len(df))
                
                # 7. 逐条吐出给 DataLoader
                for i in perm:
                    yield features_tensor[i], labels_tensor[i]
                    
        finally:
            s.close()

# ==================== 测试与使用演示 ====================
if __name__ == "__main__":
    
    # 配置你的 DolphinDB 服务器信息
    HOST = "192.168.1.100"
    PORT = 8848
    USER = "admin"
    PASSWORD = "123"
    
    DB_PATH = "dfs://HighFreqDB"
    TABLE_NAME = "tick_data"
    PARTITION_COL = "date" # 或者是你的哈希分区、范围分区的列名

    # 1. 实例化数据集
    ddb_dataset = DolphinDBPartitionForDataset(
        host=HOST, 
        port=PORT, 
        username=USER, 
        password=PASSWORD, 
        db_path=DB_PATH,
        table_name=TABLE_NAME,
        partition_col=PARTITION_COL
    )

    # 2. 扔进 DataLoader
    data_loader = DataLoader(
        dataset=ddb_dataset, 
        batch_size=256,       
        num_workers=2,        # 2个进程分别去执行不同分区的 for 循环查询
        pin_memory=False      
    )

    # 3. 测试读取
    print("开始尝试读取数据...")
    for batch_idx, (batch_x, batch_y) in enumerate(data_loader):
        print(f"\n--- 成功获取第 {batch_idx + 1} 个 Batch ---")
        print(f"X 形状: {batch_x.shape}, Y 形状: {batch_y.shape}")
        
        if batch_idx >= 2:
            print("\n基础 For 循环分区流式拉取功能测试通过！")
            break