from pathlib import PurePosixPath from typing import Tuple from airflow.providers.postgres.operators.postgres import PostgresOperator from airflow.providers.snowflake.operators.snowflake import SnowflakeOperator from airflow.utils.task_group import TaskGroup from common.operators.transform_aggregates import TransformAggregatesOperator __all__ = ["AggregatesTaskManager"] class AggregatesTaskManager: def __init__(self, keys: Tuple[str, ...], files_count: int = 100, batches_count: int = 10, transform_pool: str = "transform_pool", transform_pool_slots: int = 1): self.keys = keys self.files_count = files_count self.batches_count = batches_count self.transform_pool = transform_pool self.transform_pool_slots = transform_pool_slots self.parts = ("aggregates",) + keys self.path = "_".join(self.parts) self.batch_size = files_count // batches_count # TODO: change suffix, remove prefix self.table_name = f"test_{self.path}_airflow" self.table_name_alt = f"{self.table_name}_alt" self.data_path = PurePosixPath(*self.parts) self.raw_data_path = self.data_path / "raw_data" self.agg_data_path = self.data_path / "agg_data" def get_process_task(self): with TaskGroup(self.parts[-1]) as process: extract = SnowflakeOperator( task_id=f"extract", sql=f"{self.data_path}/extract.sql", params={ "files_count": self.files_count, "raw_data_path": self.raw_data_path }, session_parameters={"timezone": "UTC"} ) recreate_table = PostgresOperator( task_id="recreate_table", sql="postgres/recreate_table.sql", params={ "dst_table": self.table_name_alt } ) with TaskGroup("transform_and_load") as transform_and_load: for batch_index in range(self.batches_count): agg_batch_data_path = str(self.agg_data_path / f"data_{batch_index}.csv") transform = TransformAggregatesOperator( task_id=f"transform_{batch_index}", batch_index=batch_index, batch_size=self.batch_size, raw_data_path=self.raw_data_path, agg_data_path=agg_batch_data_path, pool=self.transform_pool, pool_slots=self.transform_pool_slots ) load = PostgresOperator( task_id=f"load_{batch_index}", sql="transfers/s3_to_pg.sql", params={ "s3_data_path": agg_batch_data_path, "dst_table": self.table_name_alt } ) transform >> load extract >> recreate_table >> transform_and_load return process def get_switch_task(self): switch = PostgresOperator( task_id=self.path, sql="postgres/switch_data_source.sql", params={ "src_table": self.table_name_alt, "dst_table": self.table_name } ) return switch