import argparse import datetime import importlib import logging import os from pathlib import Path from sys import argv from tadas.platform.metrics import DurationMetrics from tadas.pipeline import delete_tables, inference, publish from tadas.pipeline.date_selection import find_latest_available_date from tadas.platform import context as contexts from tadas.cli import _argparse as argparse_utils from tadas.platform import config from tadas.platform import logging as logging_utils from tadas.platform import sentry logger = logging.getLogger(__name__) MODELS_DIR = Path(__file__).resolve().parent.parent / 'models' def _discover_model_versions() -> list[str]: return sorted( path.name for path in MODELS_DIR.iterdir() if path.is_dir() and (path / 'model_config.py').is_file() ) def _create_argparse_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser() parser.add_argument( "--model", default=os.environ.get('MODEL_VERSION') or None, choices=_discover_model_versions(), help="Model version. Defaults to MODEL_VERSION env var. Must match a tadas/models subdirectory with a model_config.py", ) subparsers = parser.add_subparsers(dest="command", required=True) argparse_utils.add_inference_command(subparsers) argparse_utils.add_delete_tables_command(subparsers) argparse_utils.add_sync_to_snowflake_command(subparsers) return parser def main(args=None): parser = _create_argparse_parser() args = parser.parse_args(args=args) logger.info(args) if not args.model: parser.error("--model is required (or set MODEL_VERSION env var)") available_models = _discover_model_versions() if args.model not in available_models: parser.error( f"--model={args.model!r} is not a valid choice (got from MODEL_VERSION env). " f"Valid: {available_models}" ) model_config = importlib.import_module(f"tadas.models.{args.model}.model_config") report = args.report command = args.command context = dict( report=report, model_version=model_config.MODEL_VERSION, cli_command=command, ) with contexts.add_context(contex=context): metrics = DurationMetrics(category='cli_command_timings', context=contexts.load_context()) context_id = contexts.get_context_id(model_version=model_config.MODEL_VERSION, report=report) sentry.set_context(feed_name=context_id) logging_utils.set_context(app_context=context_id) metrics.start() try: _dispatch(command=command, args=args, model_config=model_config, report=report) finally: metrics.finish() metrics.send() def _dispatch(command, args, model_config, report): if command == 'inference': report_date = args.report_date env_report_date_str = config.get('REPORT_DATE') if not report_date and env_report_date_str: report_date = datetime.datetime.strptime(env_report_date_str, '%Y-%m-%d').date() if not report_date: report_date = find_latest_available_date( dsps=config.get('REQUIRED_DSPS') or model_config.REQUIRED_DSPS, ) inference.inference( report=report, report_date=report_date, model_config=model_config, ) elif command == 'sync-to-snowflake': publish.publish_to_snowflake(report=report, model_config=model_config) elif command == 'delete-tables': delete_tables.delete_report_tables(report=report, model_config=model_config) else: raise ValueError(f"Unknown command: {command}") if __name__ == '__main__': logging_utils.init_logging() sentry.init_sentry() main(args=argv[1:])