import os import tarfile from contextlib import contextmanager from pathlib import Path from typing import Optional, Iterable import jinja2 from airflow import AirflowException from airflow.operators.python import PythonVirtualenvOperator from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook from airflow.utils.process_utils import execute_in_subprocess from airflow.utils.python_virtualenv import prepare_virtualenv def fake_callable(): pass @contextmanager def set_cwd(path: Path): origin = Path().absolute() try: os.chdir(path) yield finally: os.chdir(origin) class DbtOperator(PythonVirtualenvOperator): def __init__( self, *, project: str, command: str, requirements: Optional[Iterable[str]] = None, conn_id: Optional[str] = 'snowflake_default', **kwargs): if not requirements: requirements = ['dbt-snowflake'] project_archive_file = Path(os.environ['AIRFLOW_HOME']) / 'dags' / f'{project}.tgz' self.project_archive_file = project_archive_file self.project = project self.conn_id = conn_id self.command = command super().__init__(python_callable=fake_callable, requirements=requirements, **kwargs) def expand_template(self, template_file: str, jinja_context: dict): template_loader = jinja2.FileSystemLoader(searchpath=os.path.dirname(__file__)) template_env = jinja2.Environment(loader=template_loader, undefined=jinja2.StrictUndefined) template = template_env.get_template(template_file) return template.render(**jinja_context) def generate_profiles_yml(self, project_dir: Path): hook = SnowflakeHook( snowlake_conn_id=self.conn_id, ) profiles_yml_ = project_dir / 'profiles.yml' conn = hook._get_conn_params() content = self.expand_template( template_file='profiles.yml.jinja2', jinja_context={'conn': conn}, ) profiles_yml_.write_text(content) def execute_callable(self): if not self.project_archive_file.exists(): raise AirflowException(f'Project archive {self.project_archive_file} not found') virtenv_dir = Path('/tmp') / f'venv_{self.project}' if self.templates_dict: self.op_kwargs['templates_dict'] = self.templates_dict prepare_virtualenv( venv_directory=str(virtenv_dir), python_bin=f'python{self.python_version}' if self.python_version else None, system_site_packages=self.system_site_packages, requirements=self.requirements, ) projects_root = Path(os.environ['AIRFLOW_HOME']) project_dir = projects_root / self.project self.log.info(f'Unpacking {self.project_archive_file} to {projects_root}') with tarfile.open(self.project_archive_file) as tar: tar.extractall(projects_root) assert project_dir.exists(), f"Archive should have created dir {project_dir}" self.generate_profiles_yml(project_dir) with set_cwd(project_dir): execute_in_subprocess( cmd=[ f'{virtenv_dir}/bin/dbt', self.command, ] )