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

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

from src.utils.database import invoke_postgres_lambda, query_snowflake
from .constants import LANGUAGE_CODES, Language

TRENDING_GROUP = "trending"


class TrendingClipAsset(ABC):
    """Base class for trending song assets."""

    def __init__(
        self,
        name: str,
        description: str,
        time_period: str,
        table_name: str,
        schedule: Optional[str] = None,
        query_file: Optional[str] = None,
    ):
        self.name = name
        self.description = description
        self.time_period = time_period
        self.table_name = table_name
        self.schedule = schedule
        self.query_file = query_file or "queries/trending_clips_template.sql"

        # 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

    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:
            query_template = f.read()

        language_codes_str = ", ".join(f"'{code}'" for code in LANGUAGE_CODES)
        query = query_template.replace("%(table_name)s", self.table_name)
        query = query.replace("%(language_codes)s", language_codes_str)
        return query

    def get_query_params(self, context: OpExecutionContext) -> Optional[Dict[str, Any]]:
        """Define parameters for the Snowflake query."""
        return None

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

        @asset(
            name=self.name,
            description=self.description,
            group_name=TRENDING_GROUP,
            tags={"tech-alerts": "true"},
        )
        def trending_clip_asset(context: OpExecutionContext) -> 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()
            query_params = instance.get_query_params(context)

            try:
                df = query_snowflake(query, query_params)
                # Log DataFrame details
                context.log.info(f"DataFrame columns: {df.columns}")
                context.log.info(f"Preview of df:\n{df.head().to_markdown()}")
                context.log.info(f"Total rows in df: {len(df)}")

            except Exception:
                context.log.error(f"Query failed with parameters: {query_params}")
                raise

            # group clips by language and create one record per language to upsert
            language_groups = df.groupby("INFERRED_LANGUAGE")
            records_to_upsert = []
            
            for language_code in Language:
                if language_code.value in language_groups.groups.keys():
                    group = language_groups.get_group(language_code.value)
                    clip_ids_list = group['CLIP_ID'].tolist()
                    records_to_upsert.append({
                        'TIME_PERIOD': instance.time_period,
                        'LANGUAGE': language_code,
                        'CLIP_IDS': clip_ids_list,
                        'DAGSTER_RUN_ID': dagster_run_id,
                    })
                else: 
                    records_to_upsert.append({
                        'TIME_PERIOD': instance.time_period,
                        'LANGUAGE': language_code,
                        'CLIP_IDS': [],
                        'DAGSTER_RUN_ID': dagster_run_id,
                    })
            
            df_to_upsert = pd.DataFrame(records_to_upsert)
            context.log.info(f"Preview of upsert dataframe: {df_to_upsert}")
            context.log.info(f"Upserting {len(df_to_upsert)} rows to Postgres")

            instance.upsert_to_postgres(df_to_upsert, context)

            return df_to_upsert

        return trending_clip_asset

    def upsert_to_postgres(self, df: pd.DataFrame, context: OpExecutionContext):
        """Shared method to handle Postgres upserts."""
        current_time = datetime.now(tz=timezone.utc)
        try:
            records = df.to_dict("records")
            upsert_stmt = """
                INSERT INTO recommendations_trendingclips
                (id, time_period, language, clip_ids, created_at, updated_at, dagster_run_id)
                VALUES (gen_random_uuid(),%s, %s, %s, %s, %s, %s)
                ON CONFLICT (time_period, language) 
                DO UPDATE SET
                    clip_ids = EXCLUDED.clip_ids,
                    updated_at = EXCLUDED.updated_at,
                    dagster_run_id = EXCLUDED.dagster_run_id
            """

            values = [
                [
                    record["TIME_PERIOD"],
                    record["LANGUAGE"],
                    record["CLIP_IDS"],
                    current_time.isoformat(),
                    current_time.isoformat(),
                    record["DAGSTER_RUN_ID"],
                ]
                for record in records
            ]

            context.log.info(f"Sending {len(values)} rows to the Postgres Lambda function")
            upsert_response = invoke_postgres_lambda(upsert_stmt, values, is_reader=False)
            context.log.info(f"Upsert response: {upsert_response}")
        except Exception as e:
            context.log.error(f"Error writing to Postgres: {e!s}")
            raise

    def create_job(self):
        """Create a job for this feature."""
        return define_asset_job(
            name=f"{self.name}_job",
            selection=AssetSelection.groups(TRENDING_GROUP),
        )

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