from __future__ import annotations
from typing import List, Optional, Tuple
from sqlalchemy import select, update, delete
from sqlalchemy.orm import Session, selectinload
from fastapi import HTTPException, status

from src.apps.feature_control.models.feature_flag import FeatureFlag
from src.apps.feature_control.models.feature_plan import FeaturePlan
from src.apps.feature_control.models.feature_plan_item import FeaturePlanItem
from src.apps.feature_control.models.merchant_feature_plan import MerchantFeaturePlan
from src.apps.feature_control.models.merchant_feature_override import MerchantFeatureOverride
from src.apps.feature_control.schemas.feature_schemas import (
    FeaturePlanCreate,
    FeaturePlanUpdate,
    FeatureOverrideItem,
)


def get_all_active_flags(db: Session) -> List[FeatureFlag]:
    stmt = select(FeatureFlag).where(FeatureFlag.is_active == True)
    return list(db.execute(stmt).scalars().all())


def get_plans(
    db: Session,
    page: int = 1,
    per_page: int = 20,
    is_active: Optional[bool] = None,
) -> Tuple[List[FeaturePlan], int]:
    stmt = (
        select(FeaturePlan)
        .where(FeaturePlan.deleted_at.is_(None))
        .options(selectinload(FeaturePlan.items))
        .order_by(FeaturePlan.created_at.desc())
    )
    if is_active is not None:
        stmt = stmt.where(FeaturePlan.is_active == is_active)

    count_stmt = stmt.with_only_columns(*[FeaturePlan.id])
    total = len(db.execute(count_stmt).all())

    stmt = stmt.offset((page - 1) * per_page).limit(per_page)
    plans = list(db.execute(stmt).scalars().unique().all())
    return plans, total


def get_plan(db: Session, plan_id: int) -> Optional[FeaturePlan]:
    stmt = (
        select(FeaturePlan)
        .where(FeaturePlan.id == plan_id, FeaturePlan.deleted_at.is_(None))
        .options(selectinload(FeaturePlan.items))
    )
    return db.execute(stmt).scalar_one_or_none()


def create_plan(db: Session, data: FeaturePlanCreate, created_by: int) -> FeaturePlan:
    plan = FeaturePlan(
        name=data.name,
        description=data.description,
        created_by=created_by,
    )
    db.add(plan)
    db.flush()

    for item_data in data.items:
        item = FeaturePlanItem(
            plan_id=plan.id,
            feature_slug=item_data.feature_slug,
            is_enabled=item_data.is_enabled,
        )
        db.add(item)

    db.flush()
    db.refresh(plan)
    return plan


def update_plan(db: Session, plan_id: int, data: FeaturePlanUpdate) -> FeaturePlan:
    plan = get_plan(db, plan_id)
    if not plan:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Feature plan not found")

    if data.name is not None:
        plan.name = data.name
    if data.description is not None:
        plan.description = data.description

    if data.items is not None:
        # Replace all items
        db.execute(
            update(FeaturePlanItem)
            .where(FeaturePlanItem.plan_id == plan_id)
            .values(is_enabled=False)
        )
        # Delete existing items and recreate
        existing = db.execute(
            select(FeaturePlanItem).where(FeaturePlanItem.plan_id == plan_id)
        ).scalars().all()
        for item in existing:
            db.delete(item)
        db.flush()

        for item_data in data.items:
            item = FeaturePlanItem(
                plan_id=plan_id,
                feature_slug=item_data.feature_slug,
                is_enabled=item_data.is_enabled,
            )
            db.add(item)

    db.flush()
    db.refresh(plan)
    return plan


def soft_delete_plan(db: Session, plan_id: int) -> FeaturePlan:
    plan = get_plan(db, plan_id)
    if not plan:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Feature plan not found")

    # Block if there are active merchant assignments
    active_count = db.execute(
        select(MerchantFeaturePlan)
        .where(MerchantFeaturePlan.plan_id == plan_id, MerchantFeaturePlan.is_active == True)
    ).scalars().first()

    if active_count:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Cannot delete plan with active merchant assignments",
        )

    from datetime import datetime, timezone
    plan.deleted_at = datetime.now(timezone.utc)
    plan.is_active = False
    db.flush()
    return plan


def get_active_plan_assignment(db: Session, merchant_id: int) -> Optional[MerchantFeaturePlan]:
    stmt = (
        select(MerchantFeaturePlan)
        .where(
            MerchantFeaturePlan.merchant_id == merchant_id,
            MerchantFeaturePlan.is_active == True,
        )
        .options(selectinload(MerchantFeaturePlan.plan).selectinload(FeaturePlan.items))
        .order_by(MerchantFeaturePlan.assigned_at.desc())
    )
    return db.execute(stmt).scalar_one_or_none()


def get_merchant_assignments(db: Session, merchant_id: int) -> List[MerchantFeaturePlan]:
    stmt = (
        select(MerchantFeaturePlan)
        .where(MerchantFeaturePlan.merchant_id == merchant_id)
        .options(selectinload(MerchantFeaturePlan.plan))
        .order_by(MerchantFeaturePlan.assigned_at.desc())
    )
    return list(db.execute(stmt).scalars().all())


def assign_plan_to_merchants(
    db: Session,
    plan_id: int,
    merchant_ids: List[int],
    assigned_by: int,
    notes: Optional[str] = None,
) -> List[int]:
    # Verify plan exists
    plan = get_plan(db, plan_id)
    if not plan:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Feature plan not found")

    affected_ids = []
    for merchant_id in merchant_ids:
        # Deactivate current assignment
        db.execute(
            update(MerchantFeaturePlan)
            .where(
                MerchantFeaturePlan.merchant_id == merchant_id,
                MerchantFeaturePlan.is_active == True,
            )
            .values(is_active=False)
        )
        assignment = MerchantFeaturePlan(
            merchant_id=merchant_id,
            plan_id=plan_id,
            assigned_by=assigned_by,
            notes=notes,
            is_active=True,
        )
        db.add(assignment)
        affected_ids.append(merchant_id)

    db.flush()
    return affected_ids


def get_overrides(db: Session, merchant_id: int) -> List[MerchantFeatureOverride]:
    stmt = select(MerchantFeatureOverride).where(
        MerchantFeatureOverride.merchant_id == merchant_id
    )
    return list(db.execute(stmt).scalars().all())


def upsert_overrides(
    db: Session,
    merchant_id: int,
    overrides: List[FeatureOverrideItem],
    set_by: int,
) -> List[MerchantFeatureOverride]:
    from datetime import datetime, timezone

    for item in overrides:
        existing = db.execute(
            select(MerchantFeatureOverride).where(
                MerchantFeatureOverride.merchant_id == merchant_id,
                MerchantFeatureOverride.feature_slug == item.feature_slug,
            )
        ).scalar_one_or_none()

        if item.is_enabled is None:
            # Clear override
            if existing:
                db.delete(existing)
        else:
            if existing:
                existing.is_enabled = item.is_enabled
                existing.reason = item.reason
                existing.set_by = set_by
                existing.set_at = datetime.now(timezone.utc)
            else:
                override = MerchantFeatureOverride(
                    merchant_id=merchant_id,
                    feature_slug=item.feature_slug,
                    is_enabled=item.is_enabled,
                    reason=item.reason,
                    set_by=set_by,
                )
                db.add(override)

    db.flush()
    return get_overrides(db, merchant_id)


def clear_overrides(db: Session, merchant_id: int) -> None:
    db.execute(
        delete(MerchantFeatureOverride).where(
            MerchantFeatureOverride.merchant_id == merchant_id
        )
    )
    db.flush()


def get_active_merchant_ids_for_plan(db: Session, plan_id: int) -> List[int]:
    rows = db.execute(
        select(MerchantFeaturePlan.merchant_id).where(
            MerchantFeaturePlan.plan_id == plan_id,
            MerchantFeaturePlan.is_active == True,
        )
    ).all()
    return [r[0] for r in rows]
