from __future__ import annotations
from datetime import datetime, timezone
from typing import List, Optional, Tuple

from fastapi import HTTPException, status
from sqlalchemy import func, select, update
from sqlalchemy.orm import Session, joinedload, selectinload

from src.apps.plan_management.models.plan import Plan
from src.apps.pricing_control.models.merchant_pricing_assignment import MerchantPricingAssignment
from src.apps.pricing_control.models.pricing_template import PricingTemplate
from src.apps.pricing_control.models.pricing_template_rate import PricingTemplateRate
from src.apps.feature_control.models.feature_plan import FeaturePlan
from src.apps.feature_control.models.feature_plan_item import FeaturePlanItem


def get_plan(db: Session, plan_id: int) -> Optional[Plan]:
    return db.execute(
        select(Plan)
        .options(
            joinedload(Plan.pricing_template).selectinload(PricingTemplate.rates),
            joinedload(Plan.feature_plan).selectinload(FeaturePlan.items),
        )
        .where(Plan.id == plan_id, Plan.deleted_at.is_(None))
    ).unique().scalar_one_or_none()


def get_plans(
    db: Session,
    page: int = 1,
    per_page: int = 20,
    is_active: Optional[bool] = None,
    search: Optional[str] = None,
) -> Tuple[List[dict], int]:
    assignment_count_sq = (
        select(
            MerchantPricingAssignment.template_id,
            func.count(MerchantPricingAssignment.id).label("cnt"),
        )
        .where(
            MerchantPricingAssignment.is_active == True,
            MerchantPricingAssignment.deleted_at.is_(None),
        )
        .group_by(MerchantPricingAssignment.template_id)
        .subquery()
    )

    q = (
        select(
            Plan.id,
            Plan.name,
            Plan.description,
            Plan.is_active,
            Plan.pricing_template_id,
            Plan.created_at,
            Plan.updated_at,
            PricingTemplate.pricing_type,
            PricingTemplate.billing_cycle,
            PricingTemplate.monthly_fee_cents,
            func.coalesce(assignment_count_sq.c.cnt, 0).label("assigned_merchant_count"),
        )
        .join(PricingTemplate, PricingTemplate.id == Plan.pricing_template_id)
        .outerjoin(assignment_count_sq, assignment_count_sq.c.template_id == Plan.pricing_template_id)
        .where(Plan.deleted_at.is_(None))
    )

    if is_active is not None:
        q = q.where(Plan.is_active == is_active)
    if search:
        q = q.where(Plan.name.ilike(f"%{search}%"))

    total = db.execute(select(func.count()).select_from(q.subquery())).scalar_one()
    rows = db.execute(q.order_by(Plan.created_at.desc()).offset((page - 1) * per_page).limit(per_page)).all()

    return [
        {
            "id": r.id,
            "name": r.name,
            "description": r.description,
            "is_active": r.is_active,
            "pricing_type": r.pricing_type,
            "billing_cycle": r.billing_cycle,
            "monthly_fee_cents": r.monthly_fee_cents,
            "created_at": r.created_at,
            "updated_at": r.updated_at,
            "assigned_merchant_count": r.assigned_merchant_count,
        }
        for r in rows
    ], total


def create_plan(
    db: Session,
    name: str,
    description: Optional[str],
    is_active: bool,
    pricing_template_id: int,
    feature_plan_id: int,
    created_by_user_id: Optional[int],
) -> Plan:
    plan = Plan(
        name=name,
        description=description,
        is_active=is_active,
        pricing_template_id=pricing_template_id,
        feature_plan_id=feature_plan_id,
        created_by_user_id=created_by_user_id,
    )
    db.add(plan)
    db.flush()
    return plan


def get_merchant_count_for_plan(db: Session, plan: Plan) -> int:
    row = db.execute(
        select(func.count(MerchantPricingAssignment.id)).where(
            MerchantPricingAssignment.template_id == plan.pricing_template_id,
            MerchantPricingAssignment.is_active == True,
            MerchantPricingAssignment.deleted_at.is_(None),
        )
    ).scalar_one()
    return row or 0


def get_affected_merchant_ids(db: Session, template_id: int) -> List[int]:
    rows = db.execute(
        select(MerchantPricingAssignment.merchant_id).where(
            MerchantPricingAssignment.template_id == template_id,
            MerchantPricingAssignment.is_active == True,
            MerchantPricingAssignment.deleted_at.is_(None),
        )
    ).all()
    return [r[0] for r in rows]


def soft_delete_plan(db: Session, plan: Plan) -> None:
    count = get_merchant_count_for_plan(db, plan)
    if count > 0:
        raise HTTPException(
            status_code=status.HTTP_409_CONFLICT,
            detail=f"Cannot delete plan with {count} active merchant assignment(s). Reassign merchants first.",
        )
    plan.deleted_at = datetime.now(timezone.utc)
    plan.is_active = False
    db.flush()


def check_plan_name_unique(db: Session, name: str, exclude_id: Optional[int] = None) -> None:
    q = select(Plan).where(Plan.name == name, Plan.deleted_at.is_(None))
    if exclude_id is not None:
        q = q.where(Plan.id != exclude_id)
    existing = db.execute(q).scalar_one_or_none()
    if existing:
        raise HTTPException(
            status_code=status.HTTP_409_CONFLICT,
            detail=f"A plan named '{name}' already exists.",
        )


def get_active_plan_for_merchant(db: Session, merchant_id: int) -> Optional[dict]:
    from src.apps.merchants.models.merchant_billing_method import MerchantBillingMethod

    row = db.execute(
        select(
            Plan.id.label("plan_id"),
            Plan.name.label("plan_name"),
            Plan.pricing_template_id,
            PricingTemplate.pricing_type,
            PricingTemplate.billing_cycle,
            Plan.feature_plan_id,
            MerchantBillingMethod.requires_revalidation,
        )
        .join(PricingTemplate, PricingTemplate.id == Plan.pricing_template_id)
        .join(MerchantPricingAssignment, MerchantPricingAssignment.template_id == Plan.pricing_template_id)
        .outerjoin(
            MerchantBillingMethod,
            (MerchantBillingMethod.merchant_id == merchant_id)
            & MerchantBillingMethod.deleted_at.is_(None)
            & (MerchantBillingMethod.is_active == True),
        )
        .where(
            MerchantPricingAssignment.merchant_id == merchant_id,
            MerchantPricingAssignment.is_active == True,
            MerchantPricingAssignment.deleted_at.is_(None),
            Plan.deleted_at.is_(None),
        )
        .order_by(MerchantPricingAssignment.effective_date.desc())
        .limit(1)
    ).first()

    if not row:
        return None

    return {
        "plan_id": row.plan_id,
        "plan_name": row.plan_name,
        "pricing_template_id": row.pricing_template_id,
        "pricing_type": row.pricing_type,
        "billing_cycle": row.billing_cycle,
        "feature_plan_id": row.feature_plan_id,
        "requires_revalidation": bool(row.requires_revalidation) if row.requires_revalidation is not None else False,
    }
