import logging
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

from src.apps.merchants.models.merchant_billing_method import MerchantBillingMethod
from src.apps.pricing_control.models.billing_record import BillingRecord
from src.apps.pricing_control.models.interchange_rate import InterchangeRate
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.pricing_control.schemas.pricing_schemas import (
    BulkAssignRequest,
    InterchangeRateCreate,
    InterchangeRateUpdate,
    MerchantBillingMethodCreate,
    MerchantPricingAssignRequest,
    PricingTemplateCreate,
    PricingTemplateUpdate,
)

logger = logging.getLogger(__name__)


# ─── Template CRUD ────────────────────────────────────────────────────────────

def get_template(db: Session, template_id: int) -> Optional[PricingTemplate]:
    return db.execute(
        select(PricingTemplate)
        .options(joinedload(PricingTemplate.rates))
        .where(PricingTemplate.id == template_id, PricingTemplate.deleted_at.is_(None))
    ).unique().scalar_one_or_none()


def get_templates(
    db: Session,
    page: int = 1,
    per_page: int = 20,
    pricing_type: Optional[str] = None,
    is_active: Optional[bool] = None,
) -> Tuple[List[PricingTemplate], int]:
    q = select(PricingTemplate).where(PricingTemplate.deleted_at.is_(None))
    if pricing_type:
        q = q.where(PricingTemplate.pricing_type == pricing_type)
    if is_active is not None:
        q = q.where(PricingTemplate.is_active == is_active)
    total = db.execute(select(func.count()).select_from(q.subquery())).scalar_one()
    items = db.execute(
        q.options(joinedload(PricingTemplate.rates))
        .order_by(PricingTemplate.created_at.desc())
        .offset((page - 1) * per_page)
        .limit(per_page)
    ).unique().scalars().all()
    return list(items), total


def create_template(
    db: Session, data: PricingTemplateCreate, created_by_user_id: Optional[int] = None
) -> PricingTemplate:
    template = PricingTemplate(
        name=data.name,
        description=data.description,
        pricing_type=data.pricing_type,
        billing_cycle=data.billing_cycle,
        monthly_fee_cents=data.monthly_fee_cents,
        setup_fee_cents=data.setup_fee_cents,
        created_by_user_id=created_by_user_id,
    )
    db.add(template)
    db.flush()

    for r in data.rates:
        rate = PricingTemplateRate(
            template_id=template.id,
            transaction_type=r.transaction_type,
            rate_percentage=r.rate_percentage,
            fixed_fee_cents=r.fixed_fee_cents,
            tier_name=r.tier_name,
            interchange_basis_points=r.interchange_basis_points,
            markup_basis_points=r.markup_basis_points,
            card_network=getattr(r, "card_network", None),
        )
        db.add(rate)

    db.commit()
    db.refresh(template)
    return template


def update_template(
    db: Session,
    template_id: int,
    data: PricingTemplateUpdate,
    updated_by_user_id: Optional[int] = None,
) -> PricingTemplate:
    old = get_template(db, template_id)
    if not old:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pricing template not found")

    # Create a new version pointing back to old
    new_template = PricingTemplate(
        name=data.name if data.name is not None else old.name,
        description=data.description if data.description is not None else old.description,
        pricing_type=old.pricing_type,
        billing_cycle=old.billing_cycle,
        monthly_fee_cents=data.monthly_fee_cents if data.monthly_fee_cents is not None else old.monthly_fee_cents,
        setup_fee_cents=data.setup_fee_cents if data.setup_fee_cents is not None else old.setup_fee_cents,
        parent_template_id=old.id,
        created_by_user_id=updated_by_user_id,
    )
    db.add(new_template)
    db.flush()

    for r in (data.rates or _copy_rates_from(db, old.id)):
        rate = PricingTemplateRate(
            template_id=new_template.id,
            transaction_type=r.transaction_type,
            rate_percentage=r.rate_percentage,
            fixed_fee_cents=r.fixed_fee_cents,
            tier_name=r.tier_name,
            interchange_basis_points=r.interchange_basis_points,
            markup_basis_points=r.markup_basis_points,
            card_network=getattr(r, "card_network", None),
        )
        db.add(rate)

    # Deactivate old template
    old.is_active = False
    db.commit()
    db.refresh(new_template)
    return new_template


def _copy_rates_from(db: Session, template_id: int):
    rates = db.execute(
        select(PricingTemplateRate).where(
            PricingTemplateRate.template_id == template_id,
            PricingTemplateRate.deleted_at.is_(None),
        )
    ).scalars().all()

    class _R:
        pass

    result = []
    for r in rates:
        obj = _R()
        obj.transaction_type = r.transaction_type
        obj.rate_percentage = r.rate_percentage
        obj.fixed_fee_cents = r.fixed_fee_cents
        obj.tier_name = r.tier_name
        obj.interchange_basis_points = r.interchange_basis_points
        obj.markup_basis_points = r.markup_basis_points
        obj.card_network = r.card_network
        result.append(obj)
    return result


def clone_template(
    db: Session,
    template_id: int,
    new_name: str,
    created_by_user_id: Optional[int] = None,
) -> PricingTemplate:
    original = get_template(db, template_id)
    if not original:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pricing template not found")

    cloned = PricingTemplate(
        name=new_name,
        description=original.description,
        pricing_type=original.pricing_type,
        billing_cycle=original.billing_cycle,
        monthly_fee_cents=original.monthly_fee_cents,
        setup_fee_cents=original.setup_fee_cents,
        parent_template_id=None,
        created_by_user_id=created_by_user_id,
    )
    db.add(cloned)
    db.flush()

    for r in _copy_rates_from(db, original.id):
        db.add(PricingTemplateRate(
            template_id=cloned.id,
            transaction_type=r.transaction_type,
            rate_percentage=r.rate_percentage,
            fixed_fee_cents=r.fixed_fee_cents,
            tier_name=r.tier_name,
            interchange_basis_points=r.interchange_basis_points,
            markup_basis_points=r.markup_basis_points,
            card_network=r.card_network,
        ))

    db.commit()
    db.refresh(cloned)
    return cloned


def soft_delete_template(db: Session, template_id: int) -> None:
    template = get_template(db, template_id)
    if not template:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pricing template not found")

    active_assignments = db.execute(
        select(MerchantPricingAssignment).where(
            MerchantPricingAssignment.template_id == template_id,
            MerchantPricingAssignment.is_active == True,
            MerchantPricingAssignment.deleted_at.is_(None),
        )
    ).scalars().first()

    if active_assignments:
        raise HTTPException(
            status_code=status.HTTP_409_CONFLICT,
            detail="Cannot delete template with active merchant assignments. Reassign merchants first.",
        )

    template.deleted_at = datetime.now(timezone.utc)
    template.is_active = False
    db.commit()


def get_template_history(db: Session, template_id: int) -> List[PricingTemplate]:
    versions: List[PricingTemplate] = []
    current_id = template_id
    visited = set()

    # Walk backwards via parent_template_id chain
    while current_id and current_id not in visited:
        visited.add(current_id)
        tpl = db.execute(
            select(PricingTemplate)
            .options(joinedload(PricingTemplate.rates))
            .where(PricingTemplate.id == current_id)
        ).unique().scalar_one_or_none()
        if not tpl:
            break
        versions.append(tpl)
        current_id = tpl.parent_template_id

    # Also collect newer versions (children pointing to this template)
    # Walk forward: find children of original requested template
    newer = db.execute(
        select(PricingTemplate)
        .options(joinedload(PricingTemplate.rates))
        .where(PricingTemplate.parent_template_id == template_id)
        .order_by(PricingTemplate.created_at.asc())
    ).unique().scalars().all()

    all_versions = list(newer) + versions
    seen_ids = set()
    unique = []
    for v in all_versions:
        if v.id not in seen_ids:
            seen_ids.add(v.id)
            unique.append(v)
    return sorted(unique, key=lambda x: x.created_at)


# ─── Assignment CRUD ──────────────────────────────────────────────────────────

def get_active_assignment(
    db: Session, merchant_id: int
) -> Optional[MerchantPricingAssignment]:
    now = datetime.now(timezone.utc)
    return db.execute(
        select(MerchantPricingAssignment)
        .options(joinedload(MerchantPricingAssignment.template))
        .where(
            MerchantPricingAssignment.merchant_id == merchant_id,
            MerchantPricingAssignment.is_active == True,
            MerchantPricingAssignment.deleted_at.is_(None),
            MerchantPricingAssignment.effective_date <= now,
        )
        .order_by(MerchantPricingAssignment.effective_date.desc())
    ).unique().scalar_one_or_none()


def get_merchant_assignments(
    db: Session, merchant_id: int
) -> List[MerchantPricingAssignment]:
    return db.execute(
        select(MerchantPricingAssignment)
        .options(joinedload(MerchantPricingAssignment.template))
        .where(
            MerchantPricingAssignment.merchant_id == merchant_id,
            MerchantPricingAssignment.deleted_at.is_(None),
        )
        .order_by(MerchantPricingAssignment.created_at.desc())
    ).unique().scalars().all()


def assign_template_to_merchants(
    db: Session,
    template_id: int,
    merchant_ids: List[int],
    effective_date: datetime,
    assigned_by_user_id: Optional[int] = None,
    notes: Optional[str] = None,
) -> List[MerchantPricingAssignment]:
    # Verify template exists
    template = get_template(db, template_id)
    if not template:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pricing template not found")

    results = []
    for merchant_id in merchant_ids:
        # Deactivate existing active assignment
        db.execute(
            update(MerchantPricingAssignment)
            .where(
                MerchantPricingAssignment.merchant_id == merchant_id,
                MerchantPricingAssignment.is_active == True,
                MerchantPricingAssignment.deleted_at.is_(None),
            )
            .values(is_active=False, updated_at=datetime.now(timezone.utc))
        )

        assignment = MerchantPricingAssignment(
            merchant_id=merchant_id,
            template_id=template_id,
            effective_date=effective_date,
            is_active=True,
            assigned_by_user_id=assigned_by_user_id,
            notes=notes,
        )
        db.add(assignment)
        results.append(assignment)

    db.flush()
    for a in results:
        db.refresh(a)
    return results


# ─── Interchange Rate CRUD ────────────────────────────────────────────────────

def get_interchange_rates(
    db: Session,
    page: int = 1,
    per_page: int = 50,
    card_network: Optional[str] = None,
    transaction_type: Optional[str] = None,
    is_active: Optional[bool] = None,
) -> Tuple[List[InterchangeRate], int]:
    q = select(InterchangeRate)
    if card_network:
        q = q.where(InterchangeRate.card_network == card_network)
    if transaction_type:
        q = q.where(InterchangeRate.transaction_type == transaction_type)
    if is_active is not None:
        q = q.where(InterchangeRate.is_active == is_active)
    total = db.execute(select(func.count()).select_from(q.subquery())).scalar_one()
    items = db.execute(
        q.order_by(InterchangeRate.card_network, InterchangeRate.card_type, InterchangeRate.transaction_type)
        .offset((page - 1) * per_page)
        .limit(per_page)
    ).scalars().all()
    return list(items), total


def create_interchange_rate(
    db: Session, data: InterchangeRateCreate, created_by: Optional[int] = None
) -> InterchangeRate:
    rate = InterchangeRate(**data.model_dump())
    db.add(rate)
    db.commit()
    db.refresh(rate)
    return rate


def update_interchange_rate(
    db: Session, rate_id: int, data: InterchangeRateUpdate
) -> InterchangeRate:
    rate = db.execute(
        select(InterchangeRate).where(InterchangeRate.id == rate_id)
    ).scalar_one_or_none()
    if not rate:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Interchange rate not found")
    for field, value in data.model_dump(exclude_none=True).items():
        setattr(rate, field, value)
    db.commit()
    db.refresh(rate)
    return rate


def deactivate_interchange_rate(db: Session, rate_id: int) -> InterchangeRate:
    rate = db.execute(
        select(InterchangeRate).where(InterchangeRate.id == rate_id)
    ).scalar_one_or_none()
    if not rate:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Interchange rate not found")
    rate.is_active = False
    rate.effective_to = datetime.now(timezone.utc)
    db.commit()
    db.refresh(rate)
    return rate


# ─── Billing Record CRUD ──────────────────────────────────────────────────────

def get_billing_records(
    db: Session,
    merchant_id: int,
    page: int = 1,
    per_page: int = 20,
) -> Tuple[List[BillingRecord], int]:
    q = select(BillingRecord).where(BillingRecord.merchant_id == merchant_id)
    total = db.execute(select(func.count()).select_from(q.subquery())).scalar_one()
    items = db.execute(
        q.order_by(BillingRecord.period_start.desc()).offset((page - 1) * per_page).limit(per_page)
    ).scalars().all()
    return list(items), total


def get_billing_record_by_period(
    db: Session, merchant_id: int, period_start: datetime
) -> Optional[BillingRecord]:
    return db.execute(
        select(BillingRecord).where(
            BillingRecord.merchant_id == merchant_id,
            BillingRecord.period_start == period_start,
        )
    ).scalar_one_or_none()


# ─── Billing Method CRUD ──────────────────────────────────────────────────────

def get_billing_method(db: Session, merchant_id: int) -> Optional[MerchantBillingMethod]:
    return db.execute(
        select(MerchantBillingMethod).where(
            MerchantBillingMethod.merchant_id == merchant_id,
            MerchantBillingMethod.is_active == True,
            MerchantBillingMethod.deleted_at.is_(None),
        )
    ).scalar_one_or_none()


def upsert_billing_method(
    db: Session, merchant_id: int, data: MerchantBillingMethodCreate
) -> MerchantBillingMethod:
    now = datetime.now(timezone.utc)
    # Soft-delete any existing method
    existing = db.execute(
        select(MerchantBillingMethod).where(
            MerchantBillingMethod.merchant_id == merchant_id,
            MerchantBillingMethod.deleted_at.is_(None),
        )
    ).scalar_one_or_none()
    if existing:
        existing.deleted_at = now
        existing.is_active = False

    method = MerchantBillingMethod(
        merchant_id=merchant_id,
        tsep_token=data.tsep_token,
        card_type=data.card_type,
        masked_card=data.masked_card,
        billing_name=data.billing_name,
        is_active=True,
    )
    db.add(method)
    db.commit()
    db.refresh(method)
    return method


# ─── Pricing summary query ────────────────────────────────────────────────────

def get_pricing_summary(
    db: Session,
    page: int = 1,
    per_page: int = 20,
    date_from: Optional[datetime] = None,
    date_to: Optional[datetime] = None,
    merchant_id: Optional[int] = None,
    pricing_type: Optional[str] = None,
) -> Tuple[List[dict], int]:
    from src.apps.merchants.models.merchant import Merchant
    from src.apps.transactions.models.transactions import Transactions
    from sqlalchemy import and_, case

    now = datetime.now(timezone.utc)

    # Subquery: get active assignment + template for each merchant
    assignment_sq = (
        select(
            MerchantPricingAssignment.merchant_id,
            MerchantPricingAssignment.template_id,
            PricingTemplate.pricing_type,
        )
        .join(PricingTemplate, PricingTemplate.id == MerchantPricingAssignment.template_id)
        .where(
            MerchantPricingAssignment.is_active == True,
            MerchantPricingAssignment.deleted_at.is_(None),
            MerchantPricingAssignment.effective_date <= now,
            PricingTemplate.deleted_at.is_(None),
        )
        .order_by(MerchantPricingAssignment.merchant_id, MerchantPricingAssignment.effective_date.desc())
        .distinct(MerchantPricingAssignment.merchant_id)
        .subquery()
    )

    txn_q = select(
        Transactions.merchant_id,
        func.sum(Transactions.txn_amount * 100).label("total_volume_cents"),
        func.sum(
            case(
                (Transactions.platform_fee_amount.isnot(None), Transactions.platform_fee_amount * 100),
                else_=0,
            )
        ).label("total_platform_fee_cents"),
    ).where(Transactions.deleted_at.is_(None))

    if date_from:
        txn_q = txn_q.where(Transactions.ocurred_at >= date_from)
    if date_to:
        txn_q = txn_q.where(Transactions.ocurred_at <= date_to)
    if merchant_id:
        txn_q = txn_q.where(Transactions.merchant_id == merchant_id)

    txn_sq = txn_q.group_by(Transactions.merchant_id).subquery()

    q = (
        select(
            Merchant.id.label("merchant_id"),
            Merchant.name.label("merchant_name"),
            assignment_sq.c.pricing_type,
            func.coalesce(txn_sq.c.total_volume_cents, 0).label("total_volume_cents"),
            func.coalesce(txn_sq.c.total_platform_fee_cents, 0).label("total_platform_fee_cents"),
        )
        .join(assignment_sq, assignment_sq.c.merchant_id == Merchant.id)
        .outerjoin(txn_sq, txn_sq.c.merchant_id == Merchant.id)
        .where(Merchant.deleted_at.is_(None))
    )

    if pricing_type:
        q = q.where(assignment_sq.c.pricing_type == pricing_type)

    total = db.execute(select(func.count()).select_from(q.subquery())).scalar_one()
    rows = db.execute(
        q.order_by(func.coalesce(txn_sq.c.total_platform_fee_cents, 0).desc())
        .offset((page - 1) * per_page)
        .limit(per_page)
    ).all()

    items = []
    for row in rows:
        vol = row.total_volume_cents or 0
        fee = row.total_platform_fee_cents or 0
        items.append({
            "merchant_id": row.merchant_id,
            "merchant_name": row.merchant_name,
            "pricing_type": row.pricing_type,
            "total_volume_cents": int(vol),
            "total_platform_fee_cents": int(fee),
            "effective_rate": round(fee / vol * 100, 4) if vol > 0 else 0.0,
        })
    return items, total
