import os
import numpy as np
from astropy.io import fits
import h5py
# from PyAstronomy import pyasl

class BaseClass():
    def _set_value(self, kwargs, key, default_value):
        ''' function to set self.key with kwargs or default value '''
        value = kwargs[key] if (key in kwargs) else default_value
        # setattr(self, key, value)
        return value

class HDF5DataSet(BaseClass):
    def __init__(self, filename):
        self.filename = filename

        if not os.path.exists(filename):
            raise FileNotFoundError(f"HDF5 file {filename} does not exist.")
        
    def read_reference_data(self):
        '''从HDF5文件中读取参考图像数据
        
        返回:
            dict: 包含image, psf, mask和header的字典
            
        异常:
            ValueError: 如果参考图像数据不存在
        '''
        with h5py.File(self.filename, 'r') as f:
            # 检查reference组是否存在
            if 'reference' not in f:
                raise ValueError("Reference data not found in the database.")
            
            reference_group = f['reference']
            
            # 检查必需的数据集是否存在
            if 'image' not in reference_group:
                raise ValueError("Reference image not found in the database.")
            if 'psf' not in reference_group:
                raise ValueError("Reference PSF not found in the database.")
            if 'mask' not in reference_group:
                raise ValueError("Reference mask not found in the database.")
            if 'header' not in reference_group:
                raise ValueError("Reference header not found in the database.")
            
            # 读取数据
            image = reference_group['image'][:]
            psf = reference_group['psf'][:]
            mask = reference_group['mask'][:]
            
            # 读取header并恢复为FITS Header对象
            header_str = reference_group['header'][()]
            if isinstance(header_str, bytes):
                header_str = header_str.decode('utf-8')
            header = fits.Header.fromstring(header_str)
            
            result = {
                'image': image,
                'psf': psf,
                'mask': mask,
                'header': header
            }
            
            return result


    def read_lc(self, star_id):
        '''从HDF5文件中读取指定星的光变曲线报告
        
        参数:
            star_id (int): 恒星ID
                
        返回:
            dict: 包含光变曲线各项指标的字典
            
        异常:
            ValueError: 如果恒星ID在数据库中不存在
        '''
        with h5py.File(self.filename, 'r') as f:
            # 检查星是否存在
            lc_path = f'lightcurve/star_{star_id}'
            if lc_path not in f:
                raise ValueError(f"star id {star_id} not exists in the database.")
            
            # 读取星的光变曲线数据
            lc_data = f[lc_path][:]
            
            # 读取图像数据并创建ID映射
            image_data = f['image/data'][:]
            image_id_to_idx = {img_id: idx for idx, img_id in enumerate(image_data['image_id'])}
            
            # 获取光变曲线中的图像ID列表
            image_ids = lc_data['image_id']
            n_lc = len(image_ids)
            
            # 为每个图像属性创建数组，使用numpy的nan作为缺失值
            attrs = ['mjd', 'fwhm', 'sky', 'qirr', 'stddev', 'airmass', 'chi2_subt', 'sigma_subt', 'npix_subt']
            image_attrs = {attr: np.full(n_lc, np.nan) for attr in attrs}
            
            # 检查 HDF5 dtype 中是否包含 DSEC 字段
            dtype_fields = image_data.dtype.names
            has_dsec = all(field in dtype_fields for field in ['dsecxmin', 'dsecxmax', 'dsecymin', 'dsecymax'])
            
            if has_dsec:
                # 如果有 DSEC 字段，添加到 attrs 中并初始化（使用 -1 作为缺失值标记，因为是整数类型）
                dsec_attrs = ['dsecxmin', 'dsecxmax', 'dsecymin', 'dsecymax']
                for attr in dsec_attrs:
                    image_attrs[attr] = np.full(n_lc, -99999, dtype='i4')
            
            # 字符串类型字段单独处理
            image_attrs['image_name'] = np.array([b''] * n_lc, dtype='S64')
            
            # 填充存在的图像数据
            for i, img_id in enumerate(image_ids):
                if img_id in image_id_to_idx:
                    idx = image_id_to_idx[img_id]
                    # 填充数值字段
                    for attr in attrs:
                        image_attrs[attr][i] = image_data[idx][attr]
                    # 填充 DSEC 字段（如果存在）
                    if has_dsec:
                        for attr in dsec_attrs:
                            image_attrs[attr][i] = image_data[idx][attr]
                    # 填充字符串字段
                    image_attrs['image_name'][i] = image_data[idx]['image_name']
            
            # 计算质量标志 (向量化)
            # flag: 0-good, 1-bad difference image, 2-bad phot, 3-both, others TBD
            bad_subt = (image_attrs['sigma_subt'] > 3.0) | (image_attrs['npix_subt'] <= 0)
            bad_phot = ((lc_data['sigma_phot'] > 3.0) | (lc_data['ferr'] <= 0) | 
                         (~np.isfinite(lc_data['flux'])) | (~np.isfinite(lc_data['ferr'])) |
                         (lc_data['npix_phot'] < 25) | (lc_data['ferr'] > 2e4))
            flags = np.zeros(len(image_ids), dtype=np.uint8)
            flags[bad_subt] = 1
            flags[bad_phot] = 2
            flags[bad_subt & bad_phot] = 3
            
            # 构建结果字典
            result = {
                'image_id': image_ids,
                'flux': lc_data['flux'],
                'eflux': lc_data['ferr'],
                'bg': lc_data['bg'],
                'ebg': lc_data['ebg'],
                'chi2_phot': lc_data['chi2_phot'],
                'sigma_phot': lc_data['sigma_phot'],
                'npix_phot': lc_data['npix_phot'],
                **{k: v for k, v in image_attrs.items()},
                'flag': flags
            }
            
            return result
        
    def read_metadata(self):
        '''读取HDF5文件中的元数据
        
        返回:
            dict: 包含元数据的字典
        '''
        std_keys = ['project', 'observatory', 'telescope', 'instrument',
                      'field_name', 'band', 'create_time', 'update_time',
                      'astrometric_reference_id', 'photometric_master_reference_id', 'photometric_reference_ids',
                      'n_images', 'n_stars']
        with h5py.File(self.filename, 'r') as f:
            meta_group = f['metadata']
            metadata = {}
            # read standard keys first
            for key in std_keys:
                if key in meta_group.attrs:
                    metadata[key] = meta_group.attrs[key]
                elif key in meta_group:
                    metadata[key] = meta_group[key][:]
            # read other keys if any
            for key in meta_group.attrs.keys():
                if key not in metadata:
                    metadata[key] = meta_group.attrs[key]
            for key in meta_group.keys():
                if key not in metadata:
                    metadata[key] = meta_group[key][:]
            
            return metadata
    
    def read_imagedata(self):
        '''读取HDF5文件中的图像数据
        
        返回:
            np.recarray: 包含图像数据的结构化数组
        '''
        with h5py.File(self.filename, 'r') as f:
            if 'image/data' not in f:
                raise ValueError(f"Image data not found in the database.")
            image_data = f['image/data'][:]
        return image_data

    def read_baseline(self, star_id=None):
        ''' 从HDF5文件中读取恒星的基线数据
        参数:
            star_id: 恒星ID，支持以下格式:
                    - None (默认): 返回所有恒星的基线数据
                    - int: 单个恒星ID
                    - list/array: 多个恒星ID的列表或数组
        返回:
            - 如果 star_id 为单个 int: 返回 dict，包含该恒星的基线数据
            - 如果 star_id 为 list/array 或 None: 返回 list of dict，每个元素包含一颗恒星的基线数据
        '''
        with h5py.File(self.filename, 'r') as f:
            # Check if baseline dataset exists
            if 'baseline/data' not in f:
                raise ValueError(f"Baseline data not found in the database.")
                
            # Read baseline data
            baseline_data = f['baseline/data'][:]
            
            # 处理不同的输入类型
            if star_id is None:
                # 返回所有恒星
                star_ids = baseline_data['star_id']
            elif isinstance(star_id, (list, np.ndarray)):
                # 输入是列表或数组
                star_ids = np.asarray(star_id)
                # 检查是否所有 ID 都存在
                missing_ids = set(star_ids) - set(baseline_data['star_id'])
                if missing_ids:
                    raise ValueError(f"Star IDs {missing_ids} not found in the baseline data.")
            else:
                # 输入是单个整数
                star_ids = np.array([star_id])
            
            # 构建结果
            results = []
            for sid in star_ids:
                mask = baseline_data['star_id'] == sid
                if not np.any(mask):
                    raise ValueError(f"Star ID {sid} not found in the baseline data.")
                
                star_data = baseline_data[mask][0]
                
                # Build result dictionary, using consistent key names with other functions
                result = {
                    'star_id': int(star_data['star_id']),
                    'x': float(star_data['x']),
                    'y': float(star_data['y']),
                    'xerr': float(getattr(star_data, 'xerr', -1.0)),
                    'yerr': float(getattr(star_data, 'yerr', -1.0)),
                    'flux': float(star_data['flux_base']),
                    'eflux': float(star_data['ferr_base']),
                    'bg': float(star_data['bg_base']),
                    'ebg': float(star_data['berr_base']),
                    'chi20': float(star_data['chi2_base'])
                }
                results.append(result)
            
            # 如果输入是单个整数，返回单个字典；否则返回列表
            if isinstance(star_id, (int, np.integer)) and not isinstance(star_id, (list, np.ndarray)):
                return results[0]
            else:
                return results

    def read_catalog_data(self):
        '''从HDF5文件中读取星表数据
        
        返回:
            dict: 以star_id为键的字典，每个值包含该星的ra, dec, x, y, mag, merr信息
            
        异常:
            ValueError: 如果星表数据不存在
        '''
        with h5py.File(self.filename, 'r') as f:
            # 检查catalog组是否存在
            if 'catalog/data' not in f:
                raise ValueError("Catalog data not found in the database.")
            
            # 读取星表数据
            catalog_data = f['catalog/data'][:]
            
            # 构建以star_id为键的字典
            result = {}
            for row in catalog_data:
                star_id = int(row['star_id'])
                result[star_id] = {
                    'ra': float(row['ra']),
                    'dec': float(row['dec']),
                    'x': float(row['x']),
                    'y': float(row['y']),
                    'mag': float(row['mag']),
                    'merr': float(row['merr'])
                }
            
            return result

