from pathlib import Path
import subprocess
import warnings
from datetime import datetime, timedelta, timezone

import dagster as dg
from dagster_dbt import (
    DbtCliResource,
    DbtProject,
    dbt_assets,
)
from dagster import (
    EnvVar,
    op,
    OpExecutionContext,
)
from dagster_snowflake import SnowflakeResource

from src.assets.dbt.analytics.dagster_dbt_translator import CustomDagsterDbtTranslator
from src.assets.dbt.analytics.factories import DbtAssetsFactory
from typing import Optional

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


def get_last_materialized_partition(
    instance: dg.DagsterInstance,
    asset_spec: dg.AssetSpec
) -> Optional[str]:
    """
    Get the last materialized partition for an asset.

    Args:
        instance: Dagster instance to query
        asset_spec: Asset spec with partitions_def

    Returns:
        The partition key of the last materialized partition, or None if:
        - Asset is not partitioned
        - No partitions have been materialized yet
    """
    asset_key = asset_spec.key
    partitions_def = asset_spec.partitions_def

    if not partitions_def:
        return None  # Asset is not partitioned

    # Get the set of materialized partition keys
    materialized_partitions = instance.get_materialized_partitions(asset_key)

    if not materialized_partitions:
        return None  # No partitions materialized yet

    # Order the materialized partitions according to the PartitionsDefinition
    # (e.g., chronologically for time-based partitions)
    all_partition_keys = partitions_def.get_partition_keys()
    ordered_materialized_partitions = [
        partition_key
        for partition_key in all_partition_keys
        if partition_key in materialized_partitions
    ]

    # The last element in the ordered list is the latest materialized partition
    return ordered_materialized_partitions[-1] if ordered_materialized_partitions else None

# Points to the dbt project path
dbt_project_directory = Path(__file__).absolute().parent / "analytics"
dbt_target = EnvVar("DBT_TARGET").get_value(default="dev")

# Validate dbt project directory exists
if not dbt_project_directory.exists():
    raise FileNotFoundError(
        f"dbt project directory not found: {dbt_project_directory}. "
        f"Please ensure the dbt project exists at this path."
    )

try:
    dbt_project = DbtProject(
        project_dir=dbt_project_directory,
        target=dbt_target
    )
    dbt_project.prepare_if_dev()
except Exception as e:
    raise RuntimeError(
        f"Failed to initialize dbt project at {dbt_project_directory}: {e}"
    ) from e

# Validate manifest exists
if not dbt_project.manifest_path.exists():
    raise FileNotFoundError(
        f"dbt manifest not found at {dbt_project.manifest_path}. "
        f"Please run 'dbt parse' or 'dbt compile' to generate the manifest."
    )

# References the dbt project object
try:
    dbt = DbtCliResource(project_dir=dbt_project_directory, target=dbt_target)
except Exception as e:
    raise RuntimeError(
        f"Failed to create DbtCliResource: {e}"
    ) from e

try:
    dbt_assets_factory = DbtAssetsFactory(
        dbt_project=dbt_project,
        dagster_dbt_translator=CustomDagsterDbtTranslator(),
    )
except Exception as e:
    raise RuntimeError(
        f"Failed to create DbtAssetsFactory: {e}"
    ) from e

# Load dbt assets with error handling
try:
    unpartitioned_dbt_models = dbt_assets_factory.create_unpartitioned_assets()
    if unpartitioned_dbt_models is None:
        raise RuntimeError("create_unpartitioned_assets() returned None instead of a list")
except Exception as e:
    raise RuntimeError(
        f"Failed to create unpartitioned dbt assets: {e}"
    ) from e

try:
    daily_partitioned_dbt_models = dbt_assets_factory.create_daily_partitioned_assets()
    if daily_partitioned_dbt_models is None:
        raise RuntimeError("create_daily_partitioned_assets() returned None instead of a list")
except Exception as e:
    raise RuntimeError(
        f"Failed to create daily partitioned dbt assets: {e}"
    ) from e

try:
    hourly_partitioned_dbt_models = dbt_assets_factory.create_hourly_partitioned_assets()
    if hourly_partitioned_dbt_models is None:
        raise RuntimeError("create_hourly_partitioned_assets() returned None instead of a list")
except Exception as e:
    raise RuntimeError(
        f"Failed to create hourly partitioned dbt assets: {e}"
    ) from e

try:
    exposure_assets = dbt_assets_factory.create_exposure_assets()
    if exposure_assets is None:
        raise RuntimeError("create_exposure_assets() returned None instead of a list")
except Exception as e:
    raise RuntimeError(
        f"Failed to create exposure assets: {e}"
    ) from e

try:
    external_source_assets = dbt_assets_factory.create_external_source_assets()
    if external_source_assets is None:
        raise RuntimeError("create_external_source_assets() returned None instead of a list")
except Exception as e:
    raise RuntimeError(
        f"Failed to create external source assets: {e}"
    ) from e

try:
    snapshot_assets = dbt_assets_factory.create_snapshot_assets()
    if snapshot_assets is None:
        raise RuntimeError("create_snapshot_assets() returned None instead of a list")
except Exception as e:
    raise RuntimeError(
        f"Failed to create snapshot assets: {e}"
    ) from e

# Log summary of loaded assets
total_assets = (
    len(unpartitioned_dbt_models) +
    len(daily_partitioned_dbt_models) +
    len(hourly_partitioned_dbt_models) +
    len(exposure_assets) +
    len(external_source_assets) +
    len(snapshot_assets)
)
print(f"✓ Successfully loaded {total_assets} dbt assets:")
print(f"  - Unpartitioned: {len(unpartitioned_dbt_models)}")
print(f"  - Daily partitioned: {len(daily_partitioned_dbt_models)}")
print(f"  - Hourly partitioned: {len(hourly_partitioned_dbt_models)}")
print(f"  - Exposures: {len(exposure_assets)}")
print(f"  - External sources: {len(external_source_assets)}")
print(f"  - Snapshots: {len(snapshot_assets)}")

# Sensor to materialize external source assets based on partition data availability
@dg.sensor(
    name="external_source_partition_availability_sensor",
    description="Sensor that checks external source assets (dbt sources without dagster.asset_key) for partition data availability and materializes partitions when data is present. Checks the last 4 partitions (hours for hourly, days for daily) and materializes any missing partitions that have data. Supports hourly and daily partition types based on source metadata configuration. Defaults to daily partitions if no partition_type is specified.",
    minimum_interval_seconds=300,  # Check every 5 minutes
)
def external_source_partition_availability_sensor(
    context: dg.SensorEvaluationContext,
    snowflake: SnowflakeResource
) -> dg.SensorResult:
    """
    Sensor that materializes external source assets based on partition data availability.

    For partitioned sources (hourly/daily):
    - Checks the last 4 partitions (hours for hourly, days for daily)
    - Queries Snowflake for missing partitions that haven't been materialized yet
    - Materializes all partitions that have data available
    - Defaults to daily partitions if no partition_type is specified
    """
    logger = context.log
    instance = context.instance

    asset_events = []
    materialized_count = 0
    skipped_count = 0
    error_count = 0

    # Get database and schema from environment or metadata
    database = EnvVar("SNOWFLAKE_DB").get_value()
    schema = EnvVar("SNOWFLAKE_SCHEMA").get_value()

    # Use a single Snowflake connection for all assets to reduce connection overhead
    # Connection is reused across all assets, with error handling to continue on failures
    try:
        with snowflake.get_connection() as conn:
            cursor = conn.cursor()

            for asset_spec in external_source_assets:
                asset_key = asset_spec.key
                partition_type = asset_spec.metadata.get("partition_type", "daily")  # Default to daily
                source_name = asset_spec.metadata.get("source_name", "")
                source_table_name = asset_spec.metadata.get("source_table_name", "")

                # Skip non-partitioned sources (partition_type=none)
                if partition_type == "none":
                    skipped_count += 1
                    logger.debug(
                        f"Skipping {asset_key.to_string()} - non-partitioned source (partition_type=none)"
                    )
                    continue

                asset_database = asset_spec.metadata.get("database") or database
                asset_schema = asset_spec.metadata.get("schema") or schema
                partition_date_column = asset_spec.metadata.get("partition_date_column", "p_date")
                partition_hour_column = asset_spec.metadata.get("partition_hour_column", "p_hour")

                # Build full table name
                table_name = f"{asset_database}.{asset_schema}.{source_table_name}"

                try:
                    # Get the actual set of materialized partitions (don't assume continuity - there may be gaps)
                    materialized_partitions = instance.get_materialized_partitions(asset_key) or set()

                    # Generate list of last 4 partitions to check
                    partitions_to_check = []
                    if partition_type == "hourly":
                        # Check last 5 hours (going backwards from previous hour to allow processing time)
                        current_hour = datetime.now(tz=timezone.utc).replace(minute=0, second=0, microsecond=0)
                        for i in range(1, 6):  # Check hours 1-5 ago
                            check_time = current_hour - timedelta(hours=i)
                            partition_key = check_time.strftime("%Y-%m-%d-%H:00")
                            partitions_to_check.append({
                                "partition_key": partition_key,
                                "partition_date": check_time.strftime("%Y-%m-%d"),
                                "partition_hour": check_time.hour,
                            })
                    else:  # daily (default)
                        # Check last 4 days (going backwards from yesterday)
                        today = datetime.now(tz=timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
                        for i in range(1, 5):  # Check days 1-4 ago
                            check_time = today - timedelta(days=i)
                            partition_key = check_time.strftime("%Y-%m-%d")
                            partitions_to_check.append({
                                "partition_key": partition_key,
                                "partition_date": check_time.strftime("%Y-%m-%d"),
                                "partition_hour": None,
                            })

                    # Filter out partitions that have already been materialized
                    missing_partitions = [
                        p for p in partitions_to_check
                        if p["partition_key"] not in materialized_partitions
                    ]

                    if not missing_partitions:
                        skipped_count += len(partitions_to_check)
                        logger.debug(
                            f"All {len(partitions_to_check)} partitions already materialized for {asset_key.to_string()}"
                        )
                        continue

                    logger.info(
                        f"Checking {len(missing_partitions)} missing partitions (out of {len(partitions_to_check)} total) "
                        f"for {asset_key.to_string()}"
                    )

                    # Check each missing partition for data availability
                    # Reuse the same cursor/connection for all assets
                    for partition_info in missing_partitions:
                        partition_key = partition_info["partition_key"]
                        partition_date = partition_info["partition_date"]
                        partition_hour = partition_info["partition_hour"]

                        try:
                            if partition_type == "hourly":
                                # Check hourly partition data with custom column names
                                query = f"""
                                SELECT COUNT(*) as total_count
                                FROM {table_name}
                                WHERE 1=1
                                    AND {partition_date_column} = '{partition_date}'
                                    AND {partition_hour_column} = {partition_hour}
                                """
                            else:  # daily
                                # Check daily partition data with custom column names
                                date_expr = f"DATE_TRUNC('day', {partition_date_column})"
                                query = f"""
                                SELECT COUNT(*) as total_count
                                FROM {table_name}
                                WHERE 1=1
                                    AND {date_expr} = '{partition_date}'
                                """

                            logger.debug(f"Executing query for {asset_key.to_string()} partition {partition_key}: {query}")
                            cursor.execute(query)
                            result = cursor.fetchone()

                            row_count = result[0] if result else 0

                            if row_count > 0:
                                # Data exists - materialize the partition
                                logger.info(
                                    f"Found {row_count} rows for {asset_key.to_string()} partition {partition_key}; "
                                    f"emitting AssetMaterialization event."
                                )

                                asset_events.append(
                                    dg.AssetMaterialization(
                                        asset_key=asset_key,
                                        partition=partition_key,
                                        description=f"External dbt source partition materialized - {row_count} rows found",
                                        metadata={
                                            "source": "external_source_sensor",
                                            "source_name": source_name,
                                            "source_table_name": source_table_name,
                                            "partition_type": partition_type,
                                            "row_count": row_count,
                                        },
                                    )
                                )
                                materialized_count += 1
                            else:
                                skipped_count += 1
                                logger.debug(
                                    f"No data found for {asset_key.to_string()} partition {partition_key}"
                                )
                        except Exception as e:
                            error_count += 1
                            logger.error(
                                f"Error checking partition {partition_key} for {asset_key.to_string()}: {str(e)}",
                                exc_info=True
                            )
                            continue

                except Exception as e:
                    error_count += 1
                    logger.error(
                        f"Error checking external source asset {asset_key.to_string()}: {str(e)}",
                        exc_info=True
                    )
                    # Continue processing other assets even if one fails
                    continue

    except Exception as e:
        # Connection-level error - log and return partial results
        error_count += 1
        logger.error(
            f"Snowflake connection error in external_source_partition_availability_sensor: {str(e)}",
            exc_info=True
        )
        # Return whatever we've collected so far

    if asset_events:
        logger.info(
            f"Materializing {materialized_count} external source asset partitions, "
            f"skipping {skipped_count} already materialized, "
            f"{error_count} errors"
        )
        return dg.SensorResult(
            asset_events=asset_events,
            cursor=str(context.cursor) if context.cursor else "initialized",
        )
    else:
        skip_reason = (
            f"Checked {len(external_source_assets)} external source assets: "
            f"{skipped_count} skipped, {error_count} errors"
        )
        logger.info(skip_reason)
        return dg.SensorResult(
            skip_reason=skip_reason,
            cursor=str(context.cursor) if context.cursor else "initialized",
        )

# Very high-frequency schedule for real-time data (every 15 minutes)
dbt_15min_schedule = dg.ScheduleDefinition(
    name="dbt_15min_models_schedule",
    cron_schedule="*/15 * * * *",  # Every 15 minutes
    job=dg.define_asset_job(
        name="dbt_15min_models_job",
        description="Run all 15min models",
        selection=dg.AssetSelection.tag("cadence", "15min"),
    ),
)

# Hourly schedule for high-frequency data (events, aggregations)
dbt_hourly_schedule = dg.ScheduleDefinition(
    name="dbt_hourly_models_schedule",
    cron_schedule="0 * * * *",  # Every hour
    job=dg.define_asset_job(
        name="dbt_hourly_models_job",
        description="Run all hourly models",
        selection=dg.AssetSelection.tag("cadence", "hourly"),
    ),
)

# Daily schedule for most transformations
dbt_daily_schedule = dg.ScheduleDefinition(
    name="dbt_daily_models_schedule",
    cron_schedule="0 2 * * *",  # Daily at 2 AM
    job=dg.define_asset_job(
        name="dbt_daily_models_job",
        description="Run all daily models",
        selection=dg.AssetSelection.tag("cadence", "daily"),
    ),
)

# Weekly schedule for mart tables and heavy computations
dbt_weekly_schedule = dg.ScheduleDefinition(
    name="dbt_weekly_models_schedule",
    cron_schedule="0 3 * * 0",  # Weekly on Sunday at 3 AM
    job=dg.define_asset_job(
        name="dbt_weekly_models_job",
        description="Run all weekly models",
        selection=dg.AssetSelection.tag("cadence", "weekly"),
    ),
)
