import os
from abc import ABC
from typing import Any, Dict, List, Optional
from datetime import datetime, timedelta, timezone

import pandas as pd
from dagster import (
    AssetKey,
    AssetSelection,
    OpExecutionContext,
    ScheduleDefinition,
    asset,
    define_asset_job,
)

from src.utils.snowflake.constants import QueryType, Role, Warehouse
from src.utils.database import execute_snowflake_ddl, query_snowflake

SNOWFLAKE_GROUP = "snowflake"


class SnowflakeJob(ABC):
    """Base class for simple Snowflake assets and jobs."""

    def __init__(
        self,
        name: str,
        description: str,
        query_file: str,
        query_type: QueryType,
        role: Optional[Role] = None,
        warehouse: Optional[Warehouse] = None,
        schedule: Optional[str] = None,
        group_name: str = SNOWFLAKE_GROUP,
        monitored: bool = True,
        owners: Optional[List[str]] = None,
        metadata: Optional[Dict[str, Any]] = None,
        tags: Optional[Dict[str, Any]] = None,
        deps: Optional[List[str]] = None,
    ):
        self.name = name
        self.description = description
        self.schedule = schedule
        self.query_file = query_file
        self.query_type = query_type
        self.role = role
        self.warehouse = warehouse
        self.group_name = group_name
        self.monitored = monitored
        self.owners = owners
        self.metadata = metadata
        self.tags = tags or {}
        self.deps = deps
        # Create asset and job
        self.asset = self.create_asset()
        self.job = self.create_job() if schedule else None
        self.schedule_def = self.create_schedule() if schedule else None

        # Merge monitored tag with existing tags if monitored is True
        if self.monitored:
            self.tags["monitored"] = "true"

    def get_query(self) -> str:
        """Load and return the SQL query from file with table name substituted."""
        query_path = os.path.join(os.path.dirname(__file__), self.query_file)
        with open(query_path, "r") as f:
            return f.read()

    def post_execute(self, result: Optional[pd.DataFrame], context: OpExecutionContext) -> None:
        """Optional hook for child classes to process query results.

        Args:
            result: DataFrame for SELECT queries, None for DDL queries
            context: Dagster execution context for logging and metadata
        """
        pass  # Default implementation does nothing

    def get_query_params(self, context: OpExecutionContext) -> Dict[str, Any]:
        """Define parameters for the Snowflake query."""
        _now = datetime.now(tz=timezone.utc)
        return {
            "start_date": _now.strftime("%Y-%m-%d"),
            "end_date": _now.strftime("%Y-%m-%d"),
            "start_hour": (_now - timedelta(hours=2)).strftime("%H"),
            "end_hour": (_now - timedelta(hours=1)).strftime("%H"),
        }

    def create_asset(self):
        """Create the Dagster asset."""
        instance = self

        @asset(
            name=self.name,
            description=self.description,
            group_name=self.group_name,
            owners=self.owners,
            metadata=self.metadata,
            tags=self.tags,
            deps=self.deps,
        )
        def snowflake_asset(context: OpExecutionContext) -> Optional[pd.DataFrame]:
            dagster_run_id = context.run_id
            context.log.info(f"DAGSTER_RUN_ID: {dagster_run_id}")

            # Get query and params
            query = instance.get_query()
            params = instance.get_query_params(context)  # Get the params

            try:
                result = None
                if self.query_type == QueryType.DDL:
                    context.log.info(
                        f"Executing DDL statement with warehouse={self.warehouse}, role={self.role}:\n{query}"
                    )
                    execute_snowflake_ddl(query, self.warehouse, self.role)
                else:
                    context.log.info(
                        f"Running select query with warehouse={self.warehouse}, role={self.role}:\n{query}, params={params}"
                    )
                    # Pass the params to query_snowflake
                    result = query_snowflake(
                        query=query, warehouse=self.warehouse, role=self.role, params=params
                    )

                # Call post_execute hook
                instance.post_execute(result, context)
                return result

            except Exception:
                context.log.error(f"Query failed")
                raise

        return snowflake_asset

    def create_job(self):
        """Create a job for this feature."""
        return define_asset_job(
            name=f"{self.name}_job",
            description=self.description,
            selection=(
                AssetSelection.groups(self.group_name).downstream()
                - AssetSelection.keys(AssetKey(["prod_marts", "active_user_metrics"]))
            ),
        )

    def create_schedule(self):
        """Create a schedule with parameters."""
        return ScheduleDefinition(
            job=self.job,
            cron_schedule=self.schedule,
            name=f"{self.name}_schedule",
        )
