"""Component for loading source files directly to staging_raw via stage.""" from datetime import datetime from garcon import task from garcon.param import StaticParam from snowflake_connector.etl_connector import SnowflakeSQLExecutor from feed_ingestion.common.staging_raw_sf.snowflake_stage_loader \ import LicensorSL from feed_ingestion.conf.config import merge_configs from feed_ingestion.flows import base from feed_ingestion.flows.deezer import config from feed_ingestion.flows.deezer.config import get_version from feed_ingestion.flows.helpers import get_sf_config from feed_ingestion.tasks import check_status class DeezerSL(LicensorSL): """StageLoader class for the Deezer ETL.""" def _fraud_db_schema(self): """Return db/schema for the fraud report SF destination.""" return { 'db': config.fraud_report_sf['db'], 'schema': config.fraud_report_sf['schema'], } def create_stage(self, stage_name, s3_dir_path, aws, **kwargs): """Create Snowflake stage. Args: stage_name (str): Name of Snowflake stage containing source files. s3_dir_path (str): s3 path to directory containing source files. aws (dict): aws credentials. """ spec_version = get_version( config.spec_version[kwargs['licensor']], datetime.strptime(stage_name[-8:], '%Y%m%d').date() ) licensor = 'all' if spec_version == 3 else kwargs['licensor'] query = 'create_snowflake_stage_{licensor}_v{version}'.format( licensor=licensor, version=spec_version) aws_params = self.get_aws_params() params = dict( db=self.executor.sf_config['db'], schema=self.executor.sf_config['schema'], stage=stage_name, s3_dir_path=s3_dir_path, **aws_params ) if 'date_format' in kwargs: params.update(date_format=kwargs['date_format']) self.resolve_sql_loader_and_execute( query, params=params ) if spec_version == 3: fraud_params = dict(params) fraud_params.update(self._fraud_db_schema()) self.resolve_sql_loader_and_execute( query + '_fraud', params=fraud_params ) def drop_stage(self, stage_name, **kwargs): """Drop Snowflake stage. Args: stage_name (str): Name of Snowflake stage containing source files. """ spec_version = get_version( config.spec_version[stage_name.split('_')[2]], datetime.strptime(stage_name[-8:], '%Y%m%d').date() ) self.resolve_sql_loader_and_execute( 'drop_snowflake_stage', params=dict( db=self.executor.sf_config['db'], schema=self.executor.sf_config['schema'], stage=stage_name ) ) if spec_version == 3: self.resolve_sql_loader_and_execute( 'drop_snowflake_stage', params=dict( **self._fraud_db_schema(), stage=stage_name + '_fraud' ) ) def clean_staging_raw_table(self, staging_raw_table, date, **kwargs): """Delete rows from previous unsuccessful workflow run. Args: staging_raw_table (str): A table name in Snowflake. date (str): Date of the data being process (YYYY-MM-DD). """ if isinstance(staging_raw_table, str): super().clean_staging_raw_table(staging_raw_table, date, **kwargs) else: fraud_table = config.snowflake_table_names['staging_raw']['v3'][-1] for table in staging_raw_table: db_schema = ( self._fraud_db_schema() if table == fraud_table else { 'db': self.executor.sf_config['db'], 'schema': self.executor.sf_config['schema'], } ) self.resolve_sql_loader_and_execute( 'delete_from_staging_raw_v3', params=dict( **db_schema, staging_raw_table=table, date=date, **kwargs ) ) def load_staging_raw_table( self, staging_raw_table, source_files_dict, date, stage_name, skip_corrupted_rows=False, **kwargs): """Load the temp_staging_raw data to the feed's staging_raw table. Args: staging_raw_table (str): A table name in Snowflake. source_files_dict (dict): A dict with source files metadata. date (str): Date of the data being process (YYYY-MM-DD). stage_name (str): Name of Snowflake stage containing source files. skip_corrupted_rows (bool): If True, add ON_ERROR=CONTINUE. """ ingestion_time = datetime.now().replace(microsecond=0) licensor = kwargs['licensor'] spec_version = get_version( config.spec_version[licensor], datetime.strptime(date, '%Y-%m-%d').date() ) query_name = 'load_staging_raw_{licensor}_v{version}'.format( licensor=licensor, version=spec_version) staging_raw_table_name = staging_raw_table if spec_version != 3 else '' stage_name_final = stage_name if spec_version != 3 else '' index = 0 for filename in source_files_dict['files']: if spec_version == 3: # remove sony_ and .csv from filename query_name = 'load_staging_raw_v3_' + filename[5:-4] staging_raw_table_name = \ config.snowflake_table_names['staging_raw']['v3'][index] if filename in config.v3_files_fraud: stage_name_final = stage_name + '_fraud' query_name = 'load_staging_raw_v3_fraud_report' staging_raw_table_name = \ config.snowflake_table_names[ 'staging_raw']['v3'][-1] fraud_name = config.fraud_report_licensor_names.get( licensor, licensor) filename = filename.replace( '*', fraud_name + '-' + date.replace('-', '') ) file_db = config.fraud_report_sf['db'] file_schema = config.fraud_report_sf['schema'] else: stage_name_final = stage_name file_db = self.executor.sf_config['db'] file_schema = self.executor.sf_config['schema'] else: file_db = self.executor.sf_config['db'] file_schema = self.executor.sf_config['schema'] params = dict( db=file_db, schema=file_schema, stage=stage_name_final, staging_raw_table=staging_raw_table_name, file_name=filename, download_date=date, ingestion_time=ingestion_time, **kwargs) sql_template = self.sql_loader.load_query(query_name) # workaround for is_valid_identifier denying '.' in identifiers # as filename having '.' in extension sql_template = sql_template.replace('%(file_name)i', filename) sql, non_identifier_params = self.executor.validator.\ format_identifiers(sql_template, params) self.executor.execute(sql, params=non_identifier_params) index += 1 @classmethod def load_activity( cls, feed_name, requirements, sql_loader=None, executor_class=SnowflakeSQLExecutor, secrets_path=None): """Load from stage activity. Args: feed_name (str): name of the feed. requirements (dict): activity requirements dict. sql_loader (SQLLoader): sql loader instance. executor_class: Snowflake executor class. secrets_path (str): Secrets manager path of the flow. """ @task.decorate(timeout=7200) @check_status(task_id='load_staging_raw_table') def load_task( activity, feed_name, date, s3_dir_path, staging_raw_table_name, source_files_dict, sfdb_params, aws, skip_corrupted_rows=False, licensor=None, report=None, date_format=None): """Copy data from Snowflake stage to staging_raw table. Args: activity (ActivityWorker): The activity worker. feed_name (str): name of the feed. date (str): Reporting date (YYYY-MM-DD). s3_dir_path (str): Name of the feed to get executor class. staging_raw_table_name (str): name of staging raw table. source_files_dict (dict): A dict with source files metadata. sfdb_params (dict): Dict with params to optionally override default ones (Snowflake db and schema name). aws (dict): aws credentials. skip_corrupted_rows (bool): If True, add ON_ERROR=CONTINUE. licensor (str): Optional licensor name. report (str): report name set query_name for load_staging_raw task. When defined then query_name == 'load_staging_raw_{report}.sql'. if not specified then query_name == 'load_staging_raw.sql' date_format (str): date format for date """ activity.logger.info('Loading staging raw table: %s', date) sf_config = get_sf_config(secrets_path) sf_config_custom = merge_configs(sf_config, sfdb_params) with executor_class(sf_config_custom) as executor: stage_loader = cls(executor, sql_loader) kwargs = { 'skip_corrupted_rows': skip_corrupted_rows, 'licensor': licensor } stage_loader.clean_staging_raw_table( staging_raw_table_name, date, **kwargs) stage_name = '{feed_name}_stage_{date:%Y%m%d}'.format( feed_name=feed_name, date=datetime.strptime(date, '%Y-%m-%d')) if date_format: kwargs.update(date_format=date_format) stage_loader.create_stage( stage_name, s3_dir_path, aws, **kwargs) if date_format: kwargs.pop('date_format') args = [ staging_raw_table_name, source_files_dict, date, stage_name] if report: kwargs['query_name'] = f'load_staging_raw_{report}' stage_loader.load_staging_raw_table(*args, **kwargs) stage_loader.drop_stage(stage_name, **kwargs) return property( lambda flow_self: flow_self.create( name='load_staging_raw_table_from_stage', tasks=base.SyncRunner( load_task.fill( aws=StaticParam(flow_self.conf_aws), **requirements) ) ) )