"""Asset checks for clip_genre_classification data quality."""

import dagster as dg
from dagster_snowflake import SnowflakeResource

from src.utils.snowflake.constants import Warehouse
from src.utils.snowflake.logger import log_query
from src.utils.snowflake.query import load_query
from src.utils.snowflake.partition_utils import parse_hourly_partition_key


@dg.asset_check(
    asset="clip_genre_classification",
    description="Verify that the clip genre table has at least a minimum row count",
    blocking=False,
)
def check_clip_genre_row_count(
    context: dg.AssetCheckExecutionContext,
    snowflake: SnowflakeResource
) -> dg.AssetCheckResult:
    """
    Check that the clip genre table partition has data (row count > 0).
    
    This check validates that the genre classification was processed successfully
    and the CLIP_GENRE table contains at least one row.
    """
    logger = context.log
    
    # Get partition information - for asset checks, use op_execution_context
    partition_key = context.op_execution_context.partition_key
    logger.info(f"Checking clip genre row count for partition: {partition_key}")
    
    # Parse partition key (format: YYYY-MM-DD-HH:00)
    p_date, p_hour = parse_hourly_partition_key(partition_key)
    
    logger.info(f"Checking partition: p_date={p_date}, p_hour={p_hour}")
    
    with snowflake.get_connection() as conn:
        cursor = conn.cursor()
        
        # Use warehouse
        warehouse_query = load_query(
            "src/utils/snowflake/queries/use_warehouse.sql",
            params={"warehouse": Warehouse.SUNO_PROD_RDS_HOURLY_X_SMALL.value}
        )
        cursor.execute(warehouse_query)
        
        # Check row count for CLIP_GENRE table
        row_count_query = f"""
            SELECT COUNT(*) as row_count
            FROM SUNO_PROD.PROD.CLIP_GENRE
            WHERE p_date = '{p_date}'
              AND p_hour = {p_hour}
        """
        
        log_query(logger, row_count_query)
        cursor.execute(row_count_query)
        result = cursor.fetchone()
        
        row_count = result[0] if result else 0
        
        logger.info(f"CLIP_GENRE row count: {row_count}")
        
        # Pass if row count > 0 (require at least 1 row)
        passed = row_count > 0
        
        return dg.AssetCheckResult(
            passed=passed,
            description=f"CLIP_GENRE partition has {row_count} rows",
            metadata={
                "partition_key": partition_key,
                "p_date": p_date,
                "p_hour": p_hour,
                "row_count": row_count,
            }
        )


# Export all asset checks
asset_checks = [
    check_clip_genre_row_count,
]

__all__ = ["asset_checks"]

