import numpy as np
import h5py
import healpy as hp
from astropy.coordinates import SkyCoord
import astropy.units as u
import os

class CatalogDatabase:
    def __init__(self, catalog_file, nside=1024):
        """
        参数:
            catalog_file: HDF5 文件路径
            nside: HEALPix 分辨率 (1024对应~3.5角分的像素)
                   如果文件已存在,将从文件读取nside,忽略此参数
        """
        self.catalog_file = catalog_file
        
        # 如果文件已存在,从文件读取nside以保证一致性
        if os.path.exists(catalog_file):
            try:
                with h5py.File(catalog_file, 'r') as f:
                    self.nside = f.attrs['nside']
                    # print(f"从文件读取 nside={self.nside}")
            except (KeyError, OSError):
                self.nside = nside
                print(f"警告: 无法从文件读取nside,使用默认值 {nside}")
        else:
            self.nside = nside
        
    def build_index(self, catalog_data):
        """
        从原始星表构建 HEALPix 索引的 HDF5 文件
        
        catalog_data: 结构化数组,包含:
            - global_id
            - ra, dec
            - mag, merr
            - stamp1, internal_id1
            - stamp2, internal_id2
        """
        # 计算每颗星的 HEALPix 像素索引
        ra = catalog_data['ra']
        dec = catalog_data['dec']
        phi = np.deg2rad(ra)
        theta = np.deg2rad(90.0 - dec)
        pix_indices = hp.ang2pix(self.nside, theta, phi)
        
        # 按 HEALPix 像素分组
        sorted_idx = np.argsort(pix_indices)
        pix_sorted = pix_indices[sorted_idx]
        
        # 找到每个像素的起始和结束位置
        unique_pix, pix_starts = np.unique(pix_sorted, return_index=True)
        pix_ends = np.append(pix_starts[1:], len(pix_sorted))
        
        # 写入 HDF5
        with h5py.File(self.catalog_file, 'w') as f:
            # 元数据
            f.attrs['nside'] = self.nside
            f.attrs['npix'] = hp.nside2npix(self.nside)
            f.attrs['total_stars'] = len(catalog_data)
            
            # 像素索引表 (快速定位)
            f.create_dataset('pixel_indices', data=unique_pix, dtype='i8')
            f.create_dataset('pixel_starts', data=pix_starts, dtype='i8')
            f.create_dataset('pixel_ends', data=pix_ends, dtype='i8')
            
            # 排序后的星表数据
            sorted_data = catalog_data[sorted_idx]
            grp = f.create_group('stars')
            grp.create_dataset('global_id', data=sorted_data['global_id'], dtype='u4')
            grp.create_dataset('ra', data=sorted_data['ra'], dtype='f8')
            grp.create_dataset('dec', data=sorted_data['dec'], dtype='f8')
            grp.create_dataset('mag', data=sorted_data['mag'], dtype='f4')
            grp.create_dataset('merr', data=sorted_data['merr'], dtype='f4')
            grp.create_dataset('internal_id1', data=sorted_data['internal_id1'], dtype='u2')
            grp.create_dataset('internal_id2', data=sorted_data['internal_id2'], dtype='u2')
            
            # 为stamp字段创建ID映射以节省空间
            # stamp1: 获取唯一值并分配ID
            unique_stamp1, stamp1_ids = np.unique(sorted_data['stamp1'], return_inverse=True)
            f.create_dataset('stamp1_names', data=unique_stamp1.astype('S20'), dtype='S20')
            grp.create_dataset('stamp1_id', data=stamp1_ids.astype('u2'), dtype='u2')
            
            # stamp2: 获取唯一值并分配ID
            unique_stamp2, stamp2_ids = np.unique(sorted_data['stamp2'], return_inverse=True)
            f.create_dataset('stamp2_names', data=unique_stamp2.astype('S20'), dtype='S20')
            grp.create_dataset('stamp2_id', data=stamp2_ids.astype('u2'), dtype='u2')
            
            print(f"stamp1 唯一值数量: {len(unique_stamp1)}, stamp2 唯一值数量: {len(unique_stamp2)}")
            
            # 创建 global_id -> 数组位置 的映射表（用于快速ID查询）
            # 假设 global_id 从 1 开始连续编号
            max_id = sorted_data['global_id'].max()
            id_to_position = np.full(max_id, -1, dtype='i4')  # -1 表示不存在
            id_to_position[sorted_data['global_id'] - 1] = np.arange(len(sorted_data), dtype='i4')
            f.create_dataset('id_to_position', data=id_to_position, dtype='i4', compression='gzip')
        
        print(f"索引构建完成: {len(unique_pix)} 个像素包含 {len(catalog_data)} 颗星")
    
    def search_circle(self, ra_center, dec_center, radius_arcsec, sort_by='separation_arcsec'):
        """
        圆锥搜索
        
        参数:
            ra_center, dec_center: 中心坐标 (度)
            radius_arcsec: 搜索半径 (角秒)
            sort_by: 排序依据,可选值: 'global_id', 'ra', 'dec', 'mag', 'merr', 
                    'separation_arcsec', 或 None (不排序)
        
        返回:
            列表,每个元素是一颗星的字典,包含该星的所有信息
        """
        # 找到可能包含目标的 HEALPix 像素
        radius_deg = radius_arcsec / 3600.0
        vec = hp.ang2vec(np.deg2rad(90.0 - dec_center), np.deg2rad(ra_center))
        
        # inclusive=True 确保即使半径很小也能包含中心像素及邻近像素
        candidate_pix = hp.query_disc(self.nside, vec, np.deg2rad(radius_deg), inclusive=True)
        
        with h5py.File(self.catalog_file, 'r') as f:
            pix_indices = f['pixel_indices'][:]
            pix_starts = f['pixel_starts'][:]
            pix_ends = f['pixel_ends'][:]
            
            # 找到候选像素在索引表中的位置
            mask = np.isin(pix_indices, candidate_pix)
            if not mask.any():
                return self._empty_result()
            
            # 读取候选像素内的所有星
            indices_to_load = []
            for pix in candidate_pix:
                loc = np.where(pix_indices == pix)[0]
                if len(loc) > 0:
                    start = pix_starts[loc[0]]
                    end = pix_ends[loc[0]]
                    indices_to_load.extend(range(start, end))
            
            if not indices_to_load:
                return self._empty_result()
            
            # 批量读取数据
            grp = f['stars']
            ra = grp['ra'][indices_to_load]
            dec = grp['dec'][indices_to_load]
            
            # 精确距离筛选 (使用角秒)
            coord_center = SkyCoord(ra_center*u.deg, dec_center*u.deg)
            coords = SkyCoord(ra*u.deg, dec*u.deg)
            seps_arcsec = coord_center.separation(coords).arcsec
            within_radius = seps_arcsec <= radius_arcsec
            
            # 提取匹配的星
            final_indices = np.array(indices_to_load)[within_radius]
            
            # 读取stamp ID并映射到名称
            stamp1_names = f['stamp1_names'][:].astype(str)
            stamp2_names = f['stamp2_names'][:].astype(str)
            stamp1_ids = grp['stamp1_id'][final_indices]
            stamp2_ids = grp['stamp2_id'][final_indices]
            
            # 组织数据为字典格式（用于排序）
            data_dict = {
                'global_id': grp['global_id'][final_indices],
                'ra': ra[within_radius],
                'dec': dec[within_radius],
                'mag': grp['mag'][final_indices],
                'merr': grp['merr'][final_indices],
                'stamp1': stamp1_names[stamp1_ids],
                'internal_id1': grp['internal_id1'][final_indices],
                'stamp2': stamp2_names[stamp2_ids],
                'internal_id2': grp['internal_id2'][final_indices],
                'separation_arcsec': seps_arcsec[within_radius]
            }
            
            # 排序
            if sort_by is not None and sort_by in data_dict:
                sort_idx = np.argsort(data_dict[sort_by])
                for key in data_dict:
                    data_dict[key] = data_dict[key][sort_idx]
            
            # 转换为列表格式
            result = []
            for i in range(len(data_dict['global_id'])):
                result.append({
                    'global_id': int(data_dict['global_id'][i]),
                    'ra': float(data_dict['ra'][i]),
                    'dec': float(data_dict['dec'][i]),
                    'mag': float(data_dict['mag'][i]),
                    'merr': float(data_dict['merr'][i]),
                    'stamp1': data_dict['stamp1'][i],
                    'internal_id1': int(data_dict['internal_id1'][i]),
                    'stamp2': data_dict['stamp2'][i],
                    'internal_id2': int(data_dict['internal_id2'][i]),
                    'separation_arcsec': float(data_dict['separation_arcsec'][i])
                })
            
        return result
    
    def search_box(self, ra_min, ra_max, dec_min, dec_max, sort_by='global_id'):
        """
        矩形区域搜索
        
        参数:
            ra_min, ra_max: RA范围 (度)
            dec_min, dec_max: Dec范围 (度)
            sort_by: 排序依据,可选值: 'global_id', 'ra', 'dec', 'mag', 'merr',
                    或 None (不排序)
        
        返回:
            列表,每个元素是一颗星的字典,包含该星的所有信息
        """
        # 使用4个角点定义矩形多边形
        corners_ra = [ra_min, ra_max, ra_max, ra_min]
        corners_dec = [dec_min, dec_min, dec_max, dec_max]
        
        theta = np.deg2rad(90.0 - np.array(corners_dec))
        phi = np.deg2rad(np.array(corners_ra))
        vertices = hp.ang2vec(theta, phi)
        
        # inclusive=True 确保边界像素也被包含
        candidate_pix = hp.query_polygon(self.nside, vertices, inclusive=True)
        
        with h5py.File(self.catalog_file, 'r') as f:
            pix_indices = f['pixel_indices'][:]
            pix_starts = f['pixel_starts'][:]
            pix_ends = f['pixel_ends'][:]
            
            indices_to_load = []
            for pix in candidate_pix:
                loc = np.where(pix_indices == pix)[0]
                if len(loc) > 0:
                    start = pix_starts[loc[0]]
                    end = pix_ends[loc[0]]
                    indices_to_load.extend(range(start, end))
            
            if not indices_to_load:
                return self._empty_result()
            
            grp = f['stars']
            ra = grp['ra'][indices_to_load]
            dec = grp['dec'][indices_to_load]
            
            # 精确边界筛选
            in_box = (ra >= ra_min) & (ra <= ra_max) & \
                     (dec >= dec_min) & (dec <= dec_max)
            
            final_indices = np.array(indices_to_load)[in_box]
            
            # 读取stamp ID并映射到名称
            stamp1_names = f['stamp1_names'][:].astype(str)
            stamp2_names = f['stamp2_names'][:].astype(str)
            stamp1_ids = grp['stamp1_id'][final_indices]
            stamp2_ids = grp['stamp2_id'][final_indices]
            
            # 组织数据为字典格式（用于排序）
            data_dict = {
                'global_id': grp['global_id'][final_indices],
                'ra': ra[in_box],
                'dec': dec[in_box],
                'mag': grp['mag'][final_indices],
                'merr': grp['merr'][final_indices],
                'stamp1': stamp1_names[stamp1_ids],
                'internal_id1': grp['internal_id1'][final_indices],
                'stamp2': stamp2_names[stamp2_ids],
                'internal_id2': grp['internal_id2'][final_indices]
            }
            
            # 排序
            if sort_by is not None and sort_by in data_dict:
                sort_idx = np.argsort(data_dict[sort_by])
                for key in data_dict:
                    data_dict[key] = data_dict[key][sort_idx]
            
            # 转换为列表格式
            result = []
            for i in range(len(data_dict['global_id'])):
                result.append({
                    'global_id': int(data_dict['global_id'][i]),
                    'ra': float(data_dict['ra'][i]),
                    'dec': float(data_dict['dec'][i]),
                    'mag': float(data_dict['mag'][i]),
                    'merr': float(data_dict['merr'][i]),
                    'stamp1': data_dict['stamp1'][i],
                    'internal_id1': int(data_dict['internal_id1'][i]),
                    'stamp2': data_dict['stamp2'][i],
                    'internal_id2': int(data_dict['internal_id2'][i])
                })
            
        return result
    
    def search_id(self, global_ids, sort_by=None):
        """
        通过 global_id 搜索恒星（使用直接索引，O(1) 复杂度）
        
        参数:
            global_ids: 单个 ID (int) 或 ID 列表 (array-like)
            sort_by: 排序依据,可选值: 'global_id', 'ra', 'dec', 'mag', 'merr',
                    或 None (不排序，保持输入顺序，未匹配的ID返回None)
        
        返回:
            列表,每个元素是一颗星的字典（包含该星的所有信息）或 None（未匹配的ID）
        """
        # 转换为数组
        if np.isscalar(global_ids):
            global_ids = np.array([global_ids], dtype='u4')
            is_scalar = True
        else:
            global_ids = np.asarray(global_ids, dtype='u4')
            is_scalar = False
        
        with h5py.File(self.catalog_file, 'r') as f:
            id_to_position_dataset = f['id_to_position']
            max_id = len(id_to_position_dataset)
            
            # 检查每个ID的有效性和存在性
            valid_mask = (global_ids > 0) & (global_ids <= max_id)
            
            # 初始化结果列表（与输入长度相同）
            result_map = {}  # 用字典存储找到的星
            
            if valid_mask.any():
                valid_ids = global_ids[valid_mask]
                valid_indices = np.where(valid_mask)[0]
                
                # 查询HDF5位置（逐个查询以避免HDF5 fancy indexing的bug）
                hdf5_positions = np.array([id_to_position_dataset[int(vid) - 1] for vid in valid_ids], dtype='i4')
                
                # 过滤出真实存在的ID（position != -1）
                exists_mask = hdf5_positions >= 0
                
                if exists_mask.any():
                    existing_ids = valid_ids[exists_mask]
                    existing_positions = hdf5_positions[exists_mask]
                    existing_input_indices = valid_indices[exists_mask]
                    
                    # 对HDF5位置排序以满足递增读取要求
                    pos_sort_indices = np.argsort(existing_positions)
                    sorted_positions = existing_positions[pos_sort_indices]
                    sorted_input_indices = existing_input_indices[pos_sort_indices]
                    
                    # 批量读取数据
                    grp = f['stars']
                    stamp1_names = f['stamp1_names'][:].astype(str)
                    stamp2_names = f['stamp2_names'][:].astype(str)
                    stamp1_ids = grp['stamp1_id'][sorted_positions]
                    stamp2_ids = grp['stamp2_id'][sorted_positions]
                    
                    # 读取所有数据
                    read_data = {
                        'global_id': grp['global_id'][sorted_positions],
                        'ra': grp['ra'][sorted_positions],
                        'dec': grp['dec'][sorted_positions],
                        'mag': grp['mag'][sorted_positions],
                        'merr': grp['merr'][sorted_positions],
                        'stamp1': stamp1_names[stamp1_ids],
                        'internal_id1': grp['internal_id1'][sorted_positions],
                        'stamp2': stamp2_names[stamp2_ids],
                        'internal_id2': grp['internal_id2'][sorted_positions]
                    }
                    
                    # 将数据映射到原始输入位置
                    for i, input_idx in enumerate(sorted_input_indices):
                        result_map[input_idx] = {
                            'global_id': int(read_data['global_id'][i]),
                            'ra': float(read_data['ra'][i]),
                            'dec': float(read_data['dec'][i]),
                            'mag': float(read_data['mag'][i]),
                            'merr': float(read_data['merr'][i]),
                            'stamp1': read_data['stamp1'][i],
                            'internal_id1': int(read_data['internal_id1'][i]),
                            'stamp2': read_data['stamp2'][i],
                            'internal_id2': int(read_data['internal_id2'][i])
                        }
        
        # 构建最终结果列表
        if sort_by is None:
            # 保持输入顺序，未匹配的返回None
            result = [result_map.get(i, None) for i in range(len(global_ids))]
        else:
            # 排序模式：只返回找到的星，按指定字段排序
            if not result_map:
                return []
            
            found_stars = list(result_map.values())
            if sort_by in found_stars[0]:
                found_stars.sort(key=lambda x: x[sort_by])
            result = found_stars
        
        return result
    
    def _empty_result(self):
        """返回空结果"""
        return []
