from pathlib import Path
src_dir = str(Path(__file__).parent) + '/src'
cat_dir = str(Path(__file__).parent) + '/data/catalog'
from src.CatalogDatabase import CatalogDatabase
import numpy as np

def str2deg(coord_str, mode):
    """
    Convert RA and Dec from sexagesimal string to decimal degrees.
    Parameters:
        coord_str (str): The coordinate string in sexagesimal format.
        mode (str): 'ra' for Right Ascension, 'dec' for Declination.
    """
    def sexagesimal_to_decimal(s):
        parts = s.split(':')
        if len(parts) == 3:
            h, m, s = map(float, parts)
            sign = -1 if h < 0 else 1
            h = abs(h)
            return sign * (h + m/60 + s/3600)
        elif len(parts) == 2:
            h, m = map(float, parts)
            sign = -1 if h < 0 else 1
            h = abs(h)
            return sign * (h + m/60)
        else:
            return float(s)
    scale = 15 if mode == 'ra' else 1  # RA is in hours, Dec is in degrees
    return sexagesimal_to_decimal(coord_str) * scale

def parse_coord(coord_str, mode):
    """
    Auto-detect coordinate format and convert to decimal degrees.
    Supports sexagesimal (hh:mm:ss / dd:mm:ss) and decimal (dd.ddddd) formats.
    """
    if ':' in str(coord_str):
        return str2deg(str(coord_str), mode)
    else:
        return float(coord_str)


class CatalogSearch:
    """
    High-level interface for searching the DREAMS star catalog.

    Usage:
        from search import CatalogSearch
        catalog = CatalogSearch()                        # use default catalog path
        catalog = CatalogSearch('path/to/catalog.h5')   # use custom catalog path

        results = catalog.search_circle(ra, dec, radius_arcsec)
        results = catalog.search_box(ra_min, ra_max, dec_min, dec_max)
        results = catalog.search_id(global_ids)
    """

    _DEFAULT_CATALOG = str(Path(__file__).parent) + '/data/catalog/DREAMS_star_catalog_v1.h5'

    def __init__(self, catalog_file=None):
        """
        Parameters:
            catalog_file (str or None): Path to the HDF5 catalog file.
                If None, uses the default path next to this script.
        """
        if catalog_file is None:
            catalog_file = self._DEFAULT_CATALOG
        self._db = CatalogDatabase(catalog_file)

    def search_circle(self, ra_center, dec_center, radius_arcsec, sort_by='separation_arcsec'):
        """
        Cone search around (ra_center, dec_center) within radius_arcsec arcseconds.
        Coordinates can be decimal degrees or sexagesimal strings (auto-detected).
        """
        ra_center  = parse_coord(ra_center,  'ra')
        dec_center = parse_coord(dec_center, 'dec')
        return self._db.search_circle(ra_center=ra_center, dec_center=dec_center,
                                      radius_arcsec=radius_arcsec, sort_by=sort_by)

    def search_box(self, ra_min, ra_max, dec_min, dec_max, sort_by='global_id'):
        """
        Box search within the given RA/Dec bounds.
        Coordinates can be decimal degrees or sexagesimal strings (auto-detected).
        """
        ra_min  = parse_coord(ra_min,  'ra')
        ra_max  = parse_coord(ra_max,  'ra')
        dec_min = parse_coord(dec_min, 'dec')
        dec_max = parse_coord(dec_max, 'dec')
        return self._db.search_box(ra_min=ra_min, ra_max=ra_max,
                                   dec_min=dec_min, dec_max=dec_max, sort_by=sort_by)

    def search_id(self, global_ids, sort_by=None):
        """
        Search by one or more global star IDs.
        """
        if not hasattr(global_ids, '__len__'):
            global_ids = np.array([global_ids], dtype='u4')
        else:
            global_ids = np.asarray(global_ids, dtype='u4')
        return self._db.search_id(global_ids=global_ids, sort_by=sort_by)


# 使用示例
if __name__ == '__main__':
    import os
    import re
    import time
    import argparse

    class CoordArgumentParser(argparse.ArgumentParser):
        """ArgumentParser that treats negative sexagesimal coords (e.g. -30:00:09.2) as positional arguments."""
        _neg_sexagesimal = re.compile(r'^-\d+:\d+')

        def _parse_optional(self, arg_string):
            if self._neg_sexagesimal.match(arg_string):
                return None
            return super()._parse_optional(arg_string)

    parser = CoordArgumentParser(description='Catalog search tool')
    subparsers = parser.add_subparsers(dest='command', help='Search mode')
    
    # Circle search subcommand
    parser_circle = subparsers.add_parser('circle', help='Cone search by coordinates')
    parser_circle.add_argument('ra', type=str, nargs='?', default='267.63463', 
                               help='Center RA (deg or hh:mm:ss, default: 267.63463)')
    parser_circle.add_argument('dec', type=str, nargs='?', default='-30.00255', 
                               help='Center Dec (deg or dd:mm:ss, default: -30.00255)')
    parser_circle.add_argument('radius', type=float, nargs='?', default=2.0, 
                               help='Search radius (arcsec, default: 2.0)')
    parser_circle.add_argument('--sort', type=str, default='separation_arcsec', 
                               choices=['separation_arcsec', 'global_id', 'ra', 'dec', 'mag', 'merr', 'none'],
                               help='Sort by (default: separation_arcsec)')
    parser_circle.add_argument('--test', action='store_true', 
                               help='Performance test mode, loop 1000 times')

    # Box search subcommand
    parser_box = subparsers.add_parser('box', help='Box search by region')
    parser_box.add_argument('ra_min', type=str, nargs='?', default='267.634', 
                            help='Min RA (deg or hh:mm:ss, default: 267.634)')
    parser_box.add_argument('ra_max', type=str, nargs='?', default='267.635', 
                            help='Max RA (deg or hh:mm:ss, default: 267.635)')
    parser_box.add_argument('dec_min', type=str, nargs='?', default='-30.003', 
                            help='Min Dec (deg or dd:mm:ss, default: -30.003)')
    parser_box.add_argument('dec_max', type=str, nargs='?', default='-30.002', 
                            help='Max Dec (deg or dd:mm:ss, default: -30.002)')
    parser_box.add_argument('--sort', type=str, default='global_id',
                            choices=['global_id', 'ra', 'dec', 'mag', 'merr', 'none'],
                            help='Sort by (default: global_id)')
    parser_box.add_argument('--test', action='store_true', 
                            help='Performance test mode, loop 1000 times')
    
    # ID search subcommand
    parser_id = subparsers.add_parser('id', help='Search by ID')
    parser_id.add_argument('ids', type=int, nargs='*', default=[46176216], 
                           help='One or more global_id (default: 46176216)')
    parser_id.add_argument('--sort', type=str, default='none',
                           choices=['global_id', 'ra', 'dec', 'mag', 'merr', 'none'],
                           help='Sort by (default: none)')
    parser_id.add_argument('--test', action='store_true', 
                           help='Performance test mode, loop 1000 times')
    
    args = parser.parse_args()

    # 初始化
    catalog = CatalogSearch()

    # 执行搜索
    if args.command == 'circle':
        sort_by = None if args.sort == 'none' else args.sort
        # 坐标转换（CatalogSearch 内部已做，这里提前转换是为了在打印中显示转换后的度数值）
        ra  = parse_coord(args.ra,  'ra')
        dec = parse_coord(args.dec, 'dec')
        nsample = 1000 if args.test else 1
        
        if args.test:
            print(f"=== Circle search performance test (loop {nsample} times) ===")
        
        time_start = time.time()
        for i in range(nsample):
            results = catalog.search_circle(ra_center=ra,
                                            dec_center=dec,
                                            radius_arcsec=args.radius,
                                            sort_by=sort_by)
        time_end = time.time()
        
        print(f"Search center: (RA={ra}°, Dec={dec}°), radius={args.radius}″")
        if args.test:
            print(f"Found {len(results)} stars, avg time {1e3*(time_end - time_start)/nsample:.3f} ms\n")
        else:
            print(f"Found {len(results)} stars, time {1e3*(time_end - time_start):.3f} ms\n")
        
        if len(results) > 0:
            print(f"{'ID':<8} {'RA[deg]':<12} {'Dec[deg]':<12} {'z_cat':<8} {'zerr_cat':<8} {'Sep[arcsec]':<12} {'Stamp1':<15} {'InternalID1':<12} {'Stamp2':<15} {'InternalID2':<12}")
            print("-" * 126)
            for star in results:
                print(f"{star['global_id']:>8d} "
                      f"{star['ra']:<12.6f} {star['dec']:<12.6f} "
                      f"{star['mag']:<8.3f} {star['merr']:<8.3f} "
                      f"{star['separation_arcsec']:<12.3f} "
                      f"{star['stamp1']:<15} {star['internal_id1']:<12} "
                      f"{star['stamp2']:<15} {star['internal_id2']:<12}")
    
    elif args.command == 'box':
        sort_by = None if args.sort == 'none' else args.sort
        ra_min  = parse_coord(args.ra_min,  'ra')
        ra_max  = parse_coord(args.ra_max,  'ra')
        dec_min = parse_coord(args.dec_min, 'dec')
        dec_max = parse_coord(args.dec_max, 'dec')
        nsample = 1000 if args.test else 1
        
        if args.test:
            print(f"=== Box search performance test (loop {nsample} times) ===")
        
        time_start = time.time()
        for i in range(nsample):
            results = catalog.search_box(ra_min=ra_min, ra_max=ra_max,
                                         dec_min=dec_min, dec_max=dec_max,
                                         sort_by=sort_by)
        time_end = time.time()
        
        print(f"Search region: RA=[{ra_min}°, {ra_max}°], Dec=[{dec_min}°, {dec_max}°]")
        if args.test:
            print(f"Found {len(results)} stars, avg time {1e3*(time_end - time_start)/nsample:.3f} ms\n")
        else:
            print(f"Found {len(results)} stars, time {1e3*(time_end - time_start):.3f} ms\n")
        
        if len(results) > 0:
            print(f"{'ID':<8} {'RA[deg]':<12} {'Dec[deg]':<12} {'z_cat':<8} {'zerr_cat':<8} {'Stamp1':<15} {'InternalID1':<12} {'Stamp2':<15} {'InternalID2':<12}")
            print("-" * 114)
            for star in results:
                print(f"{star['global_id']:>8d} "
                      f"{star['ra']:<12.6f} {star['dec']:<12.6f} "
                      f"{star['mag']:<8.3f} {star['merr']:<8.3f} "
                      f"{star['stamp1']:<15} {star['internal_id1']:<12} "
                      f"{star['stamp2']:<15} {star['internal_id2']:<12}")

    elif args.command == 'id':
        sort_by = None if args.sort == 'none' else args.sort
        global_ids = np.array(args.ids, dtype='u4')
        nsample = 1000 if args.test else 1
        
        if args.test:
            print(f"=== ID search performance test (loop {nsample} times) ===")
            print(f"Query IDs: {global_ids}\n")
        
        time_start = time.time()
        for i in range(nsample):
            results = catalog.search_id(global_ids=global_ids, sort_by=sort_by)
        time_end = time.time()
        
        if not args.test:
            print(f"Query IDs: {global_ids}")
        if args.test:
            print(f"Found {len(results)} stars, avg time {1e3*(time_end - time_start)/nsample:.3f} ms\n")
        else:
            print(f"Found {len(results)} stars, time {1e3*(time_end - time_start):.3f} ms\n")
        
        if len(results) > 0:
            print(f"{'ID':<8} {'RA[deg]':<12} {'Dec[deg]':<12} {'z_cat':<8} {'zerr_cat':<8} {'Stamp1':<15} {'InternalID1':<12} {'Stamp2':<15} {'InternalID2':<12}")
            print("-" * 114)
            for i, star in enumerate(results):
                if star is None:
                    # Unmatched ID, display as "-"
                    print(f"{global_ids[i]:>8d} "
                          f"{'-':<12} {'-':<12} "
                          f"{'-':<8} {'-':<8} "
                          f"{'-':<15} {'-':<12} "
                          f"{'-':<15} {'-':<12}")
                else:
                    print(f"{star['global_id']:>8d} "
                          f"{star['ra']:<12.6f} {star['dec']:<12.6f} "
                          f"{star['mag']:<8.3f} {star['merr']:<8.3f} "
                          f"{star['stamp1']:<15} {star['internal_id1']:<12} "
                          f"{star['stamp2']:<15} {star['internal_id2']:<12}")
    
    else:
        parser.print_help()
