import dagster as dg from dagster_snowflake import SnowflakeResource from src.utils.snowflake.constants import TIME_WINDOW_FRESHNESS_POLICY_WARN_1H_FAIL_2H, TIME_WINDOW_FRESHNESS_POLICY_WARN_24H_FAIL_25H, Group, Warehouse from src.utils.snowflake.constants import Team, SnowflakeDB, SnowflakeSchema from src.utils.snowflake.query import JinjaSQLFormatter, PythonStringSQLFormatter from src.utils.automation_conditions import daily_cron_with_eager_historical_backfill_condition ROLLING_WAU_TABLE_NAME = "ROLLING_WAU_NEW" class RollingWauConfig(dg.Config): rolling_wau_table_name: str = ROLLING_WAU_TABLE_NAME warehouse: str = Warehouse.LARGE.value @dg.asset( name="rolling_wau", description="Rolling WAU asset", group_name=Group.AGG.value, partitions_def=dg.DailyPartitionsDefinition(start_date='2024-06-01', end_offset=0), deps=[ dg.AssetDep(["prod_marts", "active_users_daily"]), dg.AssetDep(["snowflake", "bot_hourly"]), ], backfill_policy=dg.BackfillPolicy.multi_run(max_partitions_per_run=30), owners=[Team.DATA_POD.value], metadata={ "database": SnowflakeDB.SUNO_PROD.value, "schema": SnowflakeSchema.PROD.value, "table_name": ROLLING_WAU_TABLE_NAME, }, freshness_policy=TIME_WINDOW_FRESHNESS_POLICY_WARN_24H_FAIL_25H, automation_condition=daily_cron_with_eager_historical_backfill_condition ) def rolling_wau(context: dg.AssetExecutionContext, snowflake: SnowflakeResource, config: RollingWauConfig) -> dg.MaterializeResult: run_id = context.run.run_id logger = dg.get_dagster_logger() jinja_formatter = JinjaSQLFormatter() python_formatter = PythonStringSQLFormatter() # Get partition time window for processing is_multi_partition_range = context.has_partition_key_range fetch_window = context.partition_time_window fetch_start_ts = fetch_window.start fetch_end_ts = fetch_window.end fetch_params = { "rolling_wau_table_name": config.rolling_wau_table_name, "partition_start_date": fetch_start_ts.strftime("%Y-%m-%d"), "partition_end_date": fetch_end_ts.strftime("%Y-%m-%d"), } logger.info(f"Processing rolling_wau for partition: {fetch_start_ts} to {fetch_end_ts}") logger.info(f"Fetch params: {fetch_params}") with snowflake.get_connection() as conn: cursor = conn.cursor() logger.info(f"Using warehouse suno_prod_large") warehouse_query = python_formatter.load("src/utils/snowflake/queries/use_warehouse.sql", params={"warehouse": config.warehouse}, logger=logger) cursor.execute(warehouse_query) logger.info(f"Creating table {config.rolling_wau_table_name}...") create_table_query = jinja_formatter.load("src/assets/snowflake/agg/rolling_wau/table.sql", params=fetch_params, logger=logger) cursor.execute(create_table_query) logger.info(f"Deleting existing data from {config.rolling_wau_table_name} for partition window {fetch_start_ts} to {fetch_end_ts}.") delete_query = python_formatter.load("src/utils/snowflake/queries/delete_daily_partitions.sql", params={ "delete_partition_table_name": config.rolling_wau_table_name, "partition_start_date": fetch_start_ts.strftime("%Y-%m-%d"), "partition_end_date": fetch_end_ts.strftime("%Y-%m-%d"), }, logger=logger) cursor.execute(delete_query) logger.info(f"Inserting data into {config.rolling_wau_table_name}...") insert_query = jinja_formatter.load("src/assets/snowflake/agg/rolling_wau/rolling_wau.sql", params=fetch_params, logger=logger) cursor.execute(insert_query) conn.commit() rows_inserted = cursor.rowcount logger.info(f"Successfully processed partition. Inserted {rows_inserted} rows.") return dg.MaterializeResult( metadata={ "run_id": dg.MetadataValue.text(run_id), "table_name": f"{SnowflakeDB.SUNO_PROD.value}.{SnowflakeSchema.PROD.value}.{ROLLING_WAU_TABLE_NAME}", "partition_time_window_start": dg.MetadataValue.text(fetch_start_ts.isoformat()), "partition_time_window_end": dg.MetadataValue.text(fetch_end_ts.isoformat()), "dagster/row_count": rows_inserted if not is_multi_partition_range else 0, }, )