diff --git a/astrodatadownloader/astrodatadownloader.py b/astrodatadownloader/astrodatadownloader.py index a99808b..8d67c96 100644 --- a/astrodatadownloader/astrodatadownloader.py +++ b/astrodatadownloader/astrodatadownloader.py @@ -1,19 +1,32 @@ from astroquery.mast import Observations +import numpy as np +from enum import Enum -def getStarObservations(starName: str, sequences: list[int]) -> str: +class ObservationSource(Enum): + KEPLER = 1 + K2 = 2 + TESS = 3 + +def getStarObservations(starName: str, sources: list[ObservationSource], sequences: list[list[int]] = []) -> str: obs = Observations.query_object(objectname=starName, radius="0 deg") + obsWantedFilter = obs["dataproduct_type"] == "timeseries" + if(len(sequences) > 1): - obsWantedFilter = ((obs['dataproduct_type'] == 'timeseries') & - (obs['obs_collection'] == 'TESS') & - (obs['sequence_number'] >= max(sequences)) & + obsWantedFilter &= ((obs['sequence_number'] >= max(sequences)) & (obs['sequence_number'] <= min(sequences))) elif(len(sequences) == 1): - obsWantedFilter = ((obs['dataproduct_type'] == 'timeseries') & - (obs['obs_collection'] == 'TESS') & - (obs['sequence_number'] == sequences[0])) - else: - obsWantedFilter = ((obs['dataproduct_type'] == 'timeseries') & - (obs['obs_collection'] == 'TESS')) + obsWantedFilter &= (obs['sequence_number'] == sequences[0]) + + obsSourceFilter = np.ma.MaskedArray(data=np.full(obsWantedFilter.shape, False), + mask=False, fill_value=True) + if(ObservationSource.KEPLER in sources): + obsSourceFilter |= (obs['obs_collection'] == "Kepler") + if(ObservationSource.K2 in sources): + obsSourceFilter |= (obs['obs_collection'] == "K2") + if(ObservationSource.TESS in sources): + obsSourceFilter |= (obs['obs_collection'] == "TESS") + + obsWantedFilter &= obsSourceFilter return obs[obsWantedFilter] \ No newline at end of file