from typing import Any, Sequence from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.s3 import S3Hook from airflow.providers.postgres.hooks.postgres import PostgresHook from airflow.utils.context import Context from common.config.app import S3_ARTIFACTS_BUCKET from common.macros.generic import get_run_key from common.utils.s3 import S3 from common.utils.types import S3Path __all__ = ["S3ToPostgresOperator"] class S3ToPostgresOperator(BaseOperator): template_fields: Sequence[str] = ("upload_sql", "recreate_sql") template_fields_renderers = { "upload_sql": "postgresql", "recreate_sql": "postgresql" } template_ext: Sequence[str] = (".sql",) def __init__(self, upload_sql: str, recreate_sql: str, s3_data_path: S3Path, postgres_conn_id: str = PostgresHook.default_conn_name, aws_conn_id: str = S3Hook.default_conn_name, **kwargs): super().__init__(**kwargs) self.upload_sql = upload_sql self.recreate_sql = recreate_sql self.s3_data_path = s3_data_path self.postgres_conn_id = postgres_conn_id self.aws_conn_id = aws_conn_id def execute(self, context: Context) -> Any: conn = S3Hook(aws_conn_id=self.aws_conn_id).get_conn() s3 = S3(conn, run_id=get_run_key(context["dag_run"]), bucket_name=S3_ARTIFACTS_BUCKET) pg_hook = PostgresHook(postgres_conn_id=self.postgres_conn_id) pg_hook.run(self.recreate_sql) self.s3_data_path = s3.get_prefixed_key(self.s3_data_path) for key in s3.get_keys(self.s3_data_path): self.log.info(f"Loading '{key}'") pg_hook.run(self.upload_sql, parameters={"s3_data_path": key})