"""
Tender Report Service — groups transactions by payment method (card, ach, cheque, cash).
All amounts are in cents (raw txn_amount from DB).
"""
import logging
from datetime import date, datetime
from typing import Optional, List

from sqlalchemy import select, func, case, and_
from sqlalchemy.orm import Session

from src.apps.reports.schemas.payment_methods_report import (
    PaymentMethodRow,
    PaymentMethodSubRow,
    PaymentMethodReportResponse,
    PaymentMethodSummaryItem,
    PaymentMethodSummaryResponse,
)

logger = logging.getLogger(__name__)

CARD_BRANDS = ["Visa", "Mastercard", "Discover", "Amex"]


def _base_conditions(merchant_id: int, date_from, date_to):
    from src.apps.transactions.models.transactions import Transactions
    from src.apps.payment_methods.models.payment_methods import PaymentMethod

    conds = [
        Transactions.merchant_id == merchant_id,
        PaymentMethod.deleted_at.is_(None),
    ]
    if date_from:
        conds.append(Transactions.ocurred_at >= datetime.combine(date_from, datetime.min.time()))
    if date_to:
        conds.append(Transactions.ocurred_at <= datetime.combine(date_to, datetime.max.time()))
    return conds


def _query_method(db: Session, base_conditions, pm_method: str, PAID, REFUNDED, PARTIALLY_REFUNDED) -> PaymentMethodRow:
    from src.apps.transactions.models.transactions import Transactions
    from src.apps.payment_methods.models.payment_methods import PaymentMethod

    conditions = base_conditions + [PaymentMethod.method == pm_method]
    stmt = (
        select(
            func.count(case((Transactions.txn_status == PAID, 1))).label("payment_count"),
            func.count(case((Transactions.txn_status.in_([REFUNDED, PARTIALLY_REFUNDED]), 1))).label("refund_count"),
            func.coalesce(func.sum(case((Transactions.txn_status == PAID, Transactions.txn_amount), else_=0)), 0).label("payment_amount"),
            func.coalesce(func.sum(case((Transactions.txn_status.in_([REFUNDED, PARTIALLY_REFUNDED]), Transactions.txn_amount), else_=0)), 0).label("refund_amount"),
            func.coalesce(func.sum(func.coalesce(Transactions.platform_fee_amount, 0)), 0).label("fees"),
        )
        .join(PaymentMethod, Transactions.payment_method_id == PaymentMethod.id)
        .where(and_(*conditions))
    )
    r = db.execute(stmt).one()
    pa = float(r.payment_amount or 0)
    ra = float(r.refund_amount or 0)
    fees = float(r.fees or 0)
    return PaymentMethodRow(
        method=pm_method,
        payment_count=r.payment_count or 0,
        refund_count=r.refund_count or 0,
        payment_amount=pa,
        refund_amount=ra,
        fees=fees,
        net_settlement=round(pa - ra - fees, 2),
        sub_rows=[],
    )


def get_tender_report(
    db: Session,
    merchant_id: int,
    date_from: Optional[date] = None,
    date_to: Optional[date] = None,
    method: Optional[List[str]] = None,
) -> PaymentMethodReportResponse:
    from src.apps.transactions.models.transactions import Transactions
    from src.apps.payment_methods.models.payment_methods import PaymentMethod
    from src.apps.payment_methods.models.payment_method_card_details import PaymentMethodCardDetails
    from src.core.utils.enums import TransactionStatusTypes

    PAID = TransactionStatusTypes.PAID.value
    REFUNDED = TransactionStatusTypes.REFUNDED.value
    PARTIALLY_REFUNDED = TransactionStatusTypes.PARTIALLY_REFUNDED.value

    base = _base_conditions(merchant_id, date_from, date_to)
    rows: List[PaymentMethodRow] = []

    # Card — with brand sub-rows
    if not method or "card" in method:
        card_conditions = base + [PaymentMethod.method == "card"]
        brand_stmt = (
            select(
                func.coalesce(
                    func.concat(
                        func.upper(func.substr(PaymentMethodCardDetails.brand, 1, 1)),
                        func.lower(func.substr(PaymentMethodCardDetails.brand, 2)),
                    ),
                    "Other",
                ).label("brand"),
                func.count(case((Transactions.txn_status == PAID, 1))).label("payment_count"),
                func.count(case((Transactions.txn_status.in_([REFUNDED, PARTIALLY_REFUNDED]), 1))).label("refund_count"),
                func.coalesce(func.sum(case((Transactions.txn_status == PAID, Transactions.txn_amount), else_=0)), 0).label("payment_amount"),
                func.coalesce(func.sum(case((Transactions.txn_status.in_([REFUNDED, PARTIALLY_REFUNDED]), Transactions.txn_amount), else_=0)), 0).label("refund_amount"),
                func.coalesce(func.sum(func.coalesce(Transactions.platform_fee_amount, 0)), 0).label("fees"),
            )
            .join(PaymentMethod, Transactions.payment_method_id == PaymentMethod.id)
            .outerjoin(PaymentMethodCardDetails, PaymentMethod.card_details_id == PaymentMethodCardDetails.id)
            .where(and_(*card_conditions))
            .group_by("brand")
        )
        brand_results = db.execute(brand_stmt).all()

        brand_map = {}
        for r in brand_results:
            name = r.brand or "Other"
            pa = float(r.payment_amount or 0)
            ra = float(r.refund_amount or 0)
            fees = float(r.fees or 0)
            brand_map[name] = PaymentMethodSubRow(
                brand=name,
                payment_count=r.payment_count or 0,
                refund_count=r.refund_count or 0,
                payment_amount=pa,
                refund_amount=ra,
                fees=fees,
                net_settlement=round(pa - ra - fees, 2),
            )

        sub_rows = []
        ordered_brands = CARD_BRANDS + [b for b in brand_map if b not in CARD_BRANDS and b != "Other"]
        for b in ordered_brands:
            if b in brand_map:
                sub_rows.append(brand_map[b])
            else:
                sub_rows.append(PaymentMethodSubRow(brand=b, payment_count=0, refund_count=0, payment_amount=0.0, refund_amount=0.0, fees=0.0, net_settlement=0.0))
        if "Other" in brand_map:
            sub_rows.append(brand_map["Other"])

        card_pa = sum(s.payment_amount for s in sub_rows)
        card_ra = sum(s.refund_amount for s in sub_rows)
        card_fees = sum(s.fees for s in sub_rows)
        rows.append(PaymentMethodRow(
            method="card",
            payment_count=sum(s.payment_count for s in sub_rows),
            refund_count=sum(s.refund_count for s in sub_rows),
            payment_amount=round(card_pa, 2),
            refund_amount=round(card_ra, 2),
            fees=round(card_fees, 2),
            net_settlement=round(card_pa - card_ra - card_fees, 2),
            sub_rows=sub_rows,
        ))

    # ACH
    if not method or "ach" in method:
        rows.append(_query_method(db, base, "ach", PAID, REFUNDED, PARTIALLY_REFUNDED))

    # Cheque
    if not method or "cheque" in method or "check" in method:
        rows.append(_query_method(db, base, "cheque", PAID, REFUNDED, PARTIALLY_REFUNDED))

    # Cash
    if not method or "cash" in method:
        rows.append(_query_method(db, base, "cash", PAID, REFUNDED, PARTIALLY_REFUNDED))

    return PaymentMethodReportResponse(
        rows=rows,
        total_payment_amount=round(sum(r.payment_amount for r in rows), 2),
        total_refund_amount=round(sum(r.refund_amount for r in rows), 2),
        total_net_settlement=round(sum(r.net_settlement for r in rows), 2),
    )


def get_tender_summary(
    db: Session,
    merchant_id: int,
    date_from: Optional[date] = None,
    date_to: Optional[date] = None,
) -> PaymentMethodSummaryResponse:
    report = get_tender_report(db=db, merchant_id=merchant_id, date_from=date_from, date_to=date_to)
    most_used = sorted(report.rows, key=lambda r: r.payment_count, reverse=True)
    highest_volume = sorted(report.rows, key=lambda r: r.payment_amount, reverse=True)
    return PaymentMethodSummaryResponse(
        most_used=[PaymentMethodSummaryItem(method=r.method, count=r.payment_count, amount=r.payment_amount) for r in most_used],
        highest_volume=[PaymentMethodSummaryItem(method=r.method, count=r.payment_count, amount=r.payment_amount) for r in highest_volume],
    )
