import argparse import datetime import os import re from tadas.domain import constants class CustomArgumentParser(argparse.ArgumentParser): """Custom ArgumentParser which allows setting default values from ENV.""" def __init__(self, **kwargs): self.env_to_action = {} self.context_actions = [] self.args = [] super().__init__(**kwargs) def add_arg( self, name: str, flags: list = None, env_var: str = None, default=None, required=False, help: str = None, # noqa: A002 **kwargs, ): """Add argument to the parser. It supports setting the value from an environment variable. """ if not re.match(r'^[\w_]+$', name): raise ValueError(f'Invalid argument name: {name}') if default and required: raise ValueError('Cannot have both default and required set') existing_env_var = self.env_to_action.get(env_var) if existing_env_var: raise ValueError(f'Env var "{env_var}" already used by {existing_env_var}') arg = dict( name=name, flags=flags, env_var=env_var, default=default, required=required, help=help, **kwargs, ) self.args.append(arg) if not flags: flags = [f'--{name.replace("_", "-")}'] kwargs['dest'] = name help = help or f'"{name}" argument.' # noqa: A001 if default: help += f' (Default: "{default}").' kwargs['default'] = default if env_var: help += f' Can also be set by env var "{env_var}".' kwargs['help'] = help kwargs['required'] = required action = self.add_argument(*flags, **kwargs) if env_var is not None: self.env_to_action[env_var] = action def _get_env_args(self): argv = [] for env_var, action in self.env_to_action.items(): value = os.getenv(env_var) if value is not None: flag = action.option_strings[0] argv.append(f'{flag}={value}') return argv def parse_args(self, args=None, namespace=None): assert args is not None, 'args should be set' env_vars = self._get_env_args() argv_full = env_vars + args namespace = super().parse_args(argv_full, namespace) context = {} for context_action in self.context_actions: context_action.update_context(context, namespace) if context: setattr(namespace, '_context', context) return namespace def valid_date(date_str): try: return datetime.datetime.strptime(date_str, "%Y-%m-%d").date() except ValueError: raise argparse.ArgumentTypeError(f"Invalid date format: '{date_str}'. Use YYYY-MM-DD.") def add_date_argument(parser, required=True): parser.add_argument( "--date", type=valid_date, required=required, help='report_date to process. YYYY-MM-DD format') def add_sync_to_snowflake_command(subparsers): sync_to_snowflake_parser = subparsers.add_parser("sync-to-snowflake") _add_reports_arg(sync_to_snowflake_parser) def add_delete_tables_command(subparsers): delete_tables_parser = subparsers.add_parser("delete-tables") _add_reports_arg(delete_tables_parser) def add_inference_command(subparsers): inference_parser = subparsers.add_parser("inference") _add_report_date_arg(inference_parser) _add_reports_arg(inference_parser) def _add_reports_arg(run_parser: argparse.ArgumentParser): run_parser.add_argument("--report", type=str, choices=constants.REPORTS, required=True) def _add_report_date_arg(run_parser: argparse.ArgumentParser): def _valid(s: str): try: return datetime.datetime.strptime(s, "%Y-%m-%d").date() except ValueError: raise argparse.ArgumentTypeError( f"Invalid date: {s}. Expected format YYYY-MM-DD." ) run_parser.add_argument("--report-date", "--date", type=_valid, required=False)