import json
import os
from abc import ABC
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Optional
import uuid  # Add to imports at top

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 Capability, EntityType, JobGroup, JobName, Status


class EnforcementAsset(ABC):
    """Base class for enforcement assets."""

    def __init__(
        self,
        name: JobName,
        description: str,
        entity_type: EntityType,
        capability: Capability,
        status: Status,
        reason: str,
        shadow_mode: bool = False,
        schedule: Optional[str] = None,
        query_file: Optional[str] = None,
        max_entities_per_run: int = 100,
    ):
        self.name = name
        self.description = description
        self.entity_type = entity_type.value  # Convert enum to string
        self.capability = capability.value  # Convert enum to string
        self.status = status.value  # Convert enum to string
        self.reason = reason
        self.shadow_mode = shadow_mode
        self.schedule = schedule
        self.query_file = query_file or f"queries/{name}.sql"  # Default to name if not provided
        self.max_entities_per_run = max_entities_per_run
        # 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.

        Override this method if you need custom query loading logic.
        """
        query_path = os.path.join(os.path.dirname(__file__), self.query_file)
        with open(query_path, "r") as f:
            return f.read()

    def get_query_params(self, context: OpExecutionContext) -> Dict[str, Any]:
        """Define parameters for the Snowflake query."""
        _now = datetime.now(tz=timezone.utc)
        params = {
            "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"),
        }
        context.log.debug(f"Generated query parameters: {params}")
        return params

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

        @asset(
            name=self.name.value,
            description=self.description,
            group_name=JobGroup.BOT_ENFORCEMENTS.value,
            deps=[
                "bot_status_changes",
            ],
            owners=["team:core-pod"],
            metadata={
                "slack": "#tech-anti-bots",
            },
            tags={"team": "core-pod", "monitored": "true", "tech-alerts": "true"},
        )
        def enforcement_asset(context: OpExecutionContext) -> Optional[pd.DataFrame]:
            # 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 first 10 rows:\n{df.head(10).to_markdown()}")
                context.log.info(f"Total rows in df: {len(df)}")

                # Add early return if no data
                if len(df) == 0:
                    context.log.info(
                        f"No data returned from Snowflake query for {instance.name.value}"
                    )
                    return pd.DataFrame()  # Return empty DataFrame

                # Validate required fields
                if "ENTITY_ID" not in df:
                    raise ValueError(
                        "process_data must return a dict with 'entity_id' keys. "
                        f"Got: {list(df.keys())}"
                    )

                entity_ids = df["ENTITY_ID"].tolist()

                if not entity_ids:
                    context.log.info("No entity_ids to process for enforcement")
                    return None  # Changed to return None

                if len(entity_ids) > instance.max_entities_per_run:  # Changed self to instance
                    context.log.warning(
                        f"Circuit breaker: Found {len(entity_ids)} entities for enforcement, limiting to first {instance.max_entities_per_run}"
                    )
                    entity_ids = entity_ids[: instance.max_entities_per_run]

                # Write to Postgres in batches of 10K
                BATCH_SIZE = 10000
                for i in range(0, len(entity_ids), BATCH_SIZE):
                    batch = entity_ids[i : i + BATCH_SIZE]
                    instance.upsert_to_postgres(batch, context)
                    context.log.info(
                        f"Processed batch {i // BATCH_SIZE + 1} of {(len(entity_ids) + BATCH_SIZE - 1) // BATCH_SIZE}"
                    )

                return df  # Return the DataFrame

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

        return enforcement_asset

    def create_job(self):
        """Create a job for this feature."""
        return define_asset_job(
            name=f"{self.name.value}_job",
            description=self.description,
            selection=AssetSelection.assets(self.asset.key),
        )

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

    def upsert_to_postgres(self, entity_ids: List[str], context: OpExecutionContext):
        """Shared method to handle Postgres upserts."""
        current_time = datetime.now(tz=timezone.utc)
        try:
            context.log.info(
                f"Upserting bot features to Postgres with dagster run id: {context.run_id}"
            )
            upsert_stmt = """
                INSERT INTO moderation_rules AS m (
                    id,
                    entity_type,
                    entity_id,
                    capability,
                    status,
                    reason,
                    shadow_mode,
                    created_by,
                    created_at,
                    updated_at
                )
                VALUES (
                    uuid_generate_v4(),
                    %s,
                    %s,
                    %s,
                    %s,
                    %s,
                    %s,
                    %s,
                    %s,
                    %s
                )
                ON CONFLICT (entity_type, entity_id, capability) 
                DO UPDATE SET
                    capability = EXCLUDED.capability,
                    status = EXCLUDED.status,
                    reason = EXCLUDED.reason,
                    created_by = EXCLUDED.created_by,
                    updated_at = EXCLUDED.updated_at
                WHERE m.status != 'granted'
            """

            values = [
                [
                    self.entity_type,
                    str(entity_id),
                    self.capability,
                    self.status,
                    self.reason,
                    self.shadow_mode,
                    "dagster",
                    current_time.isoformat(),
                    current_time.isoformat(),
                ]
                for entity_id in entity_ids
            ]

            context.log.info(f"Upserting {len(values)} rows to Postgres")
            context.log.info(f"Upsert statement: {upsert_stmt}")
            context.log.info(f"Values: {values}")

            response = invoke_postgres_lambda(upsert_stmt, values, is_reader=False)
            if response.get("statusCode") != 200:
                raise Exception(f"Error saving bot records to Postgres: {response.get('body')}")
            context.log.info(f"Save results:\n{json.dumps(response, indent=2)}")
        except Exception as e:
            context.log.error(f"Error writing to Postgres: {e!s}")
            raise
