
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,
        },
    )
