from datetime import datetime
import dagster as dg
from dagster import EnvVar
from dagster_snowflake import SnowflakeResource

from src.assets.snowflake.dim.user.assets import dim_user
from src.assets.snowflake.raw_table.frontend.web_audio_player_actions.assets import web_audio_player_actions
from src.utils.automation_conditions import hourly_cron_with_eager_historical_backfill_condition
from src.utils.snowflake.constants import TIME_WINDOW_FRESHNESS_POLICY_WARN_1H_FAIL_2H, PartitionExpr, Team, Warehouse
from src.utils.snowflake.query import JinjaSQLFormatter
from src.utils.snowflake.query import PythonStringSQLFormatter


FACT_STUDIO_PLAY_START_DATE = datetime.strptime('2025-09-24', '%Y-%m-%d')
FACT_STUDIO_PLAY_TABLE_NAME = "FACT_STUDIO_PLAY"

class FactStudioPlayConfig(dg.Config):
    studio_event_table_name: str = "WEB_AUDIO_PLAYER_ACTIONS"
    dedup_table_name: str = "TEMP_STUDIO_PLAYS_DEDUP"
    unique_segments_table_name: str = "TEMP_STUDIO_PLAYS_UNIQUE_SEGMENTS"
    temp_studio_project_clip_sessions_table_name: str = "TEMP_STUDIO_PROJECT_CLIP_SESSIONS"
    stg_studio_play_table_name: str = "STG_STUDIO_PLAY"
    fact_studio_play_table_name: str = FACT_STUDIO_PLAY_TABLE_NAME
    warehouse: str = Warehouse.FACT_PLAY_SMALL.value

@dg.asset(
    name="fact_studio_play",
    description="Fact table for Studio plays. One row per each clip played in a Studio project session; this means that a user could have multiple simultaneous studio play session entries if they are playing multiple tracks at once.",
    group_name="studio",
    partitions_def=dg.HourlyPartitionsDefinition(start_date=FACT_STUDIO_PLAY_START_DATE, end_offset=-3),
    deps=[
        dg.AssetDep(web_audio_player_actions, partition_mapping=dg.TimeWindowPartitionMapping(start_offset=-3, end_offset=3)),
        dg.AssetDep(["PROD", "dim_clip"]),
        dg.AssetDep(dim_user),
    ],
    backfill_policy=dg.BackfillPolicy.multi_run(max_partitions_per_run=24*7*2),
    owners=[Team.DATA_POD.value],
    metadata={
        "database": EnvVar("SNOWFLAKE_DB").get_value(),
        "schema": EnvVar("SNOWFLAKE_SCHEMA").get_value(),
        "table_name": FACT_STUDIO_PLAY_TABLE_NAME,
        "data_start_date": FACT_STUDIO_PLAY_START_DATE.strftime("%Y-%m-%d"),
        "partition_expr": PartitionExpr.HOURLY.value,
        "sla_minutes": 180
    },
    automation_condition=hourly_cron_with_eager_historical_backfill_condition,
    freshness_policy=TIME_WINDOW_FRESHNESS_POLICY_WARN_1H_FAIL_2H,
)
def fact_studio_play(context: dg.AssetExecutionContext, snowflake: SnowflakeResource, config: FactStudioPlayConfig) -> dg.MaterializeResult:
    run_id = context.run.run_id
    logger = dg.get_dagster_logger()
    jinja_formatter = JinjaSQLFormatter()
    python_formatter = PythonStringSQLFormatter()
    is_multi_partition_range = context.has_partition_key_range

    fetch_window = context.asset_partitions_time_window_for_input(web_audio_player_actions.key.to_user_string())
    fetch_start_ts = min(fetch_window.start, fetch_window.start)
    fetch_end_ts = max(fetch_window.end, fetch_window.end)

    fetch_params = {
        "buffered_partition_start_date": fetch_start_ts.strftime("%Y-%m-%d"),
        "buffered_partition_end_date": fetch_end_ts.strftime("%Y-%m-%d"),
        "buffered_partition_start_hour": fetch_start_ts.hour,
        "buffered_partition_end_hour": fetch_end_ts.hour,
        "partition_start_date": context.partition_time_window.start.strftime("%Y-%m-%d"),
        "partition_end_date": context.partition_time_window.end.strftime("%Y-%m-%d"),
        "partition_start_hour": context.partition_time_window.start.hour,
        "partition_end_hour": context.partition_time_window.end.hour,
        "studio_event_table_name": config.studio_event_table_name,
        "dedup_table_name": config.dedup_table_name,
        "unique_segments_table_name": config.unique_segments_table_name,
        "temp_studio_project_clip_sessions_table_name": config.temp_studio_project_clip_sessions_table_name,
        "stg_studio_play_table_name": config.stg_studio_play_table_name,
        "fact_studio_play_table_name": config.fact_studio_play_table_name,
    }

    logger.info(f"Processing target partition: {context.partition_time_window.start} to {context.partition_time_window.end}")
    logger.info(f"Fetching raw data from buffered window: {context.partition_time_window.start} to {context.partition_time_window.end}")

    with snowflake.get_connection() as conn:
        cursor = conn.cursor()
        logger.info(f"Using warehouse {config.warehouse}")
        warehouse_query = python_formatter.load("src/utils/snowflake/queries/use_warehouse.sql", params={"warehouse": config.warehouse}, logger=logger)
        cursor.execute(warehouse_query)

        # 1. Create temp table with deduplicated data from the buffered window for the partition(s)
        logger.info(f"Creating temporary table {config.dedup_table_name} with deduplicated Studio play session-clip entries...")
        dedup_query = jinja_formatter.load("src/assets/snowflake/fact/studio_play/queries/stg_studio_events_deduplicated.sql", params=fetch_params, logger=logger)
        cursor.execute(dedup_query)

        # 2. Create temp table with unique Studio play segments for the partition(s)
        logger.info(f"Creating temporary table {config.unique_segments_table_name} with unique play segments...")
        unique_segments_query = jinja_formatter.load("src/assets/snowflake/fact/studio_play/queries/stg_studio_unique_play_segments.sql", params=fetch_params, logger=logger)
        cursor.execute(unique_segments_query)

        # 3. Create temp table with aggregated Studio project session x clip entries
        logger.info(f"Creating temporary table {config.temp_studio_project_clip_sessions_table_name} with aggregated Studio project session-clip entries...")
        sessions_query = jinja_formatter.load("src/assets/snowflake/fact/studio_play/queries/stg_studio_project_clip_sessions.sql", params=fetch_params, logger=logger)
        cursor.execute(sessions_query)

        # 4. Create temp table with aggregated Studio project session x clip entries with misc. metadata
        logger.info(f"Creating temporary table {config.stg_studio_play_table_name} with aggregated Studio project session-clip entries with misc. metadata...")
        sessions_query = jinja_formatter.load("src/assets/snowflake/fact/studio_play/queries/stg_studio_play.sql", params=fetch_params, logger=logger)
        cursor.execute(sessions_query)

        # 4. Delete existing data from the target table for the partition(s)
        logger.info(f"Deleting existing data from {config.fact_studio_play_table_name} for partition(s) {context.partition_time_window.start} to {context.partition_time_window.end}")
        delete_query = python_formatter.load("src/utils/snowflake/queries/delete_hourly_partitions.sql", params={**fetch_params, "delete_partition_table_name": config.fact_studio_play_table_name}, logger=logger)
        cursor.execute(delete_query)

        # 5. Merge staged sessions into fact table
        logger.info(f"Inserting staged Studio play session-clip entries from {config.stg_studio_play_table_name} into final fact table {config.fact_studio_play_table_name}...")
        insert_fact_studio_play_query = jinja_formatter.load("src/assets/snowflake/fact/studio_play/queries/insert_stg_to_final_table.sql", params=fetch_params, logger=logger)
        cursor.execute(insert_fact_studio_play_query)
        rows_inserted = cursor.rowcount

    logger.info(f"Successfully processed partition(s) {context.partition_time_window.start} to {context.partition_time_window.end}. Inserted {rows_inserted} rows")

    return dg.MaterializeResult(
        metadata={
            "run_id": dg.MetadataValue.text(run_id),
            "table_name": config.fact_studio_play_table_name,
            "partition_time_window_start": dg.MetadataValue.text(context.partition_time_window.start.isoformat()),
            "partition_time_window_end": dg.MetadataValue.text(context.partition_time_window.end.isoformat()),
            "dagster/row_count": rows_inserted if not is_multi_partition_range else 0,
        },
    )
