import warnings
from datetime import datetime

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

# Using asset key references to avoid import chain issues
from src.utils.automation_conditions import hourly_cron_with_eager_historical_backfill_except_usage_plan_condition
from src.utils.snowflake.constants import TIME_WINDOW_FRESHNESS_POLICY_WARN_1H_FAIL_2H, Group, PartitionExpr, Warehouse
from src.utils.snowflake.query import JinjaSQLFormatter, PythonStringSQLFormatter

warnings.filterwarnings("ignore", category=dg.BetaWarning)

DIM_USER_START_DATE = datetime.strptime('2023-03-01', '%Y-%m-%d')
DIM_USER_TABLE_NAME = "DIM_USER"


class DimUserConfig(dg.Config):
    dim_user_table_name: str = DIM_USER_TABLE_NAME
    warehouse: str = Warehouse.LARGE.value


@dg.asset(
    name="dim_user",
    description="Dimension table containing user information including account details, subscription status, and geographic data.",
    group_name=Group.IDENTITY.value,
    partitions_def=dg.HourlyPartitionsDefinition(start_date=DIM_USER_START_DATE, end_offset=-1),
    deps=[
        dg.AssetDep("rds_discord_info"),
        dg.AssetDep("rds_auth_user"),
        dg.AssetDep("rds_usage_plan", partition_mapping=dg.AllPartitionMapping()),
        dg.AssetDep("web_user_event"),
        dg.AssetDep("app_event"),
        dg.AssetDep("user_session_state"),
    ],
    backfill_policy=dg.BackfillPolicy.multi_run(max_partitions_per_run=24*7),
    owners=["team:core-pod"],
    metadata={
        "database": EnvVar("SNOWFLAKE_DB").get_value(),
        "schema": EnvVar("SNOWFLAKE_SCHEMA").get_value(),
        "table_name": DIM_USER_TABLE_NAME,
        "data_start_date": DIM_USER_START_DATE.strftime("%Y-%m-%d"),
        "cluster_by": "[p_date, p_hour]",
        "partition_expr": PartitionExpr.HOURLY.value,
        "transient": True,
        "sla_minutes": 180,
    },
    automation_condition=hourly_cron_with_eager_historical_backfill_except_usage_plan_condition,
    freshness_policy=TIME_WINDOW_FRESHNESS_POLICY_WARN_1H_FAIL_2H,
)
def dim_user(context: dg.AssetExecutionContext, snowflake: SnowflakeResource, config: DimUserConfig) -> 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
    partition_start = context.partition_time_window.start
    partition_end = context.partition_time_window.end

    fetch_params = {
        "partition_count_filter": 'multiple' if context.has_partition_key_range else 'single',
        "partition_start_date": partition_start.strftime("%Y-%m-%d"),
        "partition_end_date": partition_end.strftime("%Y-%m-%d"),
        "partition_start_hour": partition_start.hour,
        "partition_end_hour": partition_end.hour,
        "dim_user_table_name": config.dim_user_table_name,
    }

    logger.info(f"Processing dim_user for partition: {partition_start} to {partition_end}")
    logger.info(f"Note: Processing 60 days of data (from {partition_start.strftime('%Y-%m-%d')} - 60 days to current time)")
    logger.info(f"Fetch params: {fetch_params}")

    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. Ensure the target table exists
        logger.info(f"Creating table {config.dim_user_table_name}...")
        create_table_query = jinja_formatter.load("src/assets/snowflake/dim/user/table.sql", params=fetch_params, logger=logger)
        logger.info(f"Table creation query: {create_table_query}")
        try:
            cursor.execute(create_table_query)
            logger.info("Table creation successful")
        except Exception as e:
            logger.error(f"Table creation failed: {e}")
            raise

        # 2. Delete existing data from the target table for the current partition only
        logger.info(f"Deleting existing data from {config.dim_user_table_name} for partition window {partition_start} to {partition_end}.")
        delete_query = python_formatter.load("src/utils/snowflake/queries/delete_hourly_partitions.sql", params={**fetch_params, "delete_partition_table_name": config.dim_user_table_name}, logger=logger)
        logger.info(f"Delete query: {delete_query}")
        try:
            cursor.execute(delete_query)
            logger.info("Delete query successful")
        except Exception as e:
            logger.error(f"Delete query failed: {e}")
            raise

        # 3. Insert/update user dimension data using MERGE logic
        logger.info(f"Inserting/updating user dimension data in {config.dim_user_table_name}.")
        logger.info("Note: Only users with changed attributes will be updated for existing users beyond 60 days.")
        insert_query = jinja_formatter.load("src/assets/snowflake/dim/user/insert_dim_user.sql", params=fetch_params, logger=logger)

        # Execute each MERGE statement separately but within the same transaction
        logger.info("Executing MERGE operations as separate statements within transaction")
        statements = [stmt.strip() for stmt in insert_query.split(';') if stmt.strip()]

        for i, statement in enumerate(statements):
            if statement:
                logger.info(f"Executing MERGE statement {i+1}/{len(statements)}")
                try:
                    cursor.execute(statement)
                    logger.info(f"MERGE statement {i+1} successful")
                except Exception as e:
                    logger.error(f"MERGE statement {i+1} failed: {e}")
                    raise

        # 4. Commit the transaction
        conn.commit()
        rows_affected = cursor.rowcount

    logger.info(f"Successfully processed partition. Affected {rows_affected} rows.")

    return dg.MaterializeResult(
        metadata={
            "run_id": dg.MetadataValue.text(run_id),
            "table_name": config.dim_user_table_name,
            "partition_time_window_start": dg.MetadataValue.text(partition_start.isoformat()),
            "partition_time_window_end": dg.MetadataValue.text(partition_end.isoformat()),
            "dagster/row_count": rows_affected,
        },
    )
