"""
Celery task: generate monthly platform billing records and charge merchants.

Task name: pricing_control.generate_platform_billing
Schedule:  1st of each month at 02:00 UTC (Celery Beat)

For each merchant with an active pricing assignment:
  1. Aggregate platform_fee_amount from transactions in the previous month
  2. Add monthly platform fee from the template
  3. Create a BillingRecord (idempotent — skips if record already exists)
  4. Charge via the merchant's stored billing method using HW_TSYS platform account
  5. Retry up to 3 times on charge failure
"""
import logging
from calendar import monthrange
from datetime import datetime, timezone

from celery.utils.log import get_task_logger
from sqlalchemy import func, select
from sqlalchemy.orm import joinedload

from src.worker.celery_app import celery_app
from src.core.database import SessionCelery

logger = get_task_logger(__name__)


def _get_previous_month_bounds(now: datetime):
    """Return (period_start, period_end) for the previous calendar month."""
    year = now.year
    month = now.month - 1
    if month == 0:
        month = 12
        year -= 1
    _, last_day = monthrange(year, month)
    period_start = datetime(year, month, 1, 0, 0, 0, tzinfo=timezone.utc)
    period_end = datetime(year, month, last_day, 23, 59, 59, tzinfo=timezone.utc)
    return period_start, period_end


def _is_billing_due(billing_cycle: str, now: datetime) -> bool:
    """Return True if this billing cycle should be charged in the current month."""
    month = now.month
    if billing_cycle == "monthly":
        return True
    if billing_cycle == "quarterly":
        return month in (1, 4, 7, 10)
    if billing_cycle == "yearly":
        return month == 1
    return True


@celery_app.task(name="pricing_control.generate_platform_billing", bind=True, max_retries=3)
def generate_platform_billing(self) -> dict:
    """
    Monthly task — generates billing records for all merchants with active pricing assignments
    and attempts to charge each via the merchant's stored billing method.
    """
    from src.apps.pricing_control.models.billing_record import BillingRecord, BillingStatus
    from src.apps.pricing_control.models.merchant_pricing_assignment import MerchantPricingAssignment
    from src.apps.pricing_control.crud import get_active_assignment, get_billing_record_by_period, get_billing_method
    from src.apps.transactions.models.transactions import Transactions
    from src.apps.merchants.models.merchant import Merchant

    now = datetime.now(timezone.utc)
    period_start, period_end = _get_previous_month_bounds(now)

    logger.info("generate_platform_billing: processing period %s → %s", period_start, period_end)

    created_count = 0
    skipped_count = 0
    failed_count = 0

    with SessionCelery() as db:
        # Get all distinct merchant IDs with an active pricing assignment
        merchant_ids = db.execute(
            select(MerchantPricingAssignment.merchant_id)
            .where(
                MerchantPricingAssignment.is_active == True,
                MerchantPricingAssignment.deleted_at.is_(None),
            )
            .distinct()
        ).scalars().all()

    for merchant_id in merchant_ids:
        try:
            with SessionCelery() as db:
                # Idempotency: skip if billing record already exists for this period
                existing = get_billing_record_by_period(db, merchant_id, period_start)
                if existing:
                    logger.debug("Billing record exists for merchant %s period %s — skipping", merchant_id, period_start)
                    skipped_count += 1
                    continue

                assignment = get_active_assignment(db, merchant_id)
                if not assignment:
                    logger.warning("No active assignment for merchant %s — skipping", merchant_id)
                    skipped_count += 1
                    continue

                template = assignment.template
                if not _is_billing_due(template.billing_cycle, now):
                    logger.debug("Billing not due this month for merchant %s (cycle=%s)", merchant_id, template.billing_cycle)
                    skipped_count += 1
                    continue

                # Sum platform_fee_amount for previous period (stored as dollars; convert to cents)
                fee_sum = db.execute(
                    select(func.coalesce(func.sum(Transactions.platform_fee_amount), 0.0)).where(
                        Transactions.merchant_id == merchant_id,
                        Transactions.deleted_at.is_(None),
                        Transactions.ocurred_at >= period_start,
                        Transactions.ocurred_at <= period_end,
                        Transactions.platform_fee_amount.isnot(None),
                    )
                ).scalar_one()
                txn_fee_cents = int(round(float(fee_sum or 0) * 100))

                # Merchant volume for the period
                vol_sum = db.execute(
                    select(func.coalesce(func.sum(Transactions.txn_amount), 0.0)).where(
                        Transactions.merchant_id == merchant_id,
                        Transactions.deleted_at.is_(None),
                        Transactions.ocurred_at >= period_start,
                        Transactions.ocurred_at <= period_end,
                    )
                ).scalar_one()
                volume_cents = int(round(float(vol_sum or 0) * 100))

                monthly_fee_cents = template.monthly_fee_cents or 0
                total_due = txn_fee_cents + monthly_fee_cents

                record = BillingRecord(
                    merchant_id=merchant_id,
                    template_id=template.id,
                    period_start=period_start,
                    period_end=period_end,
                    total_volume_cents=volume_cents,
                    total_platform_fee_cents=txn_fee_cents,
                    monthly_fee_cents=monthly_fee_cents,
                    total_due_cents=total_due,
                    status=BillingStatus.PENDING,
                )
                db.add(record)
                db.commit()
                db.refresh(record)
                created_count += 1

                # Attempt charge
                billing_method = get_billing_method(db, merchant_id)
                if not billing_method:
                    logger.warning("Merchant %s has no billing method — skipping charge", merchant_id)
                    record.status = BillingStatus.SKIPPED
                    db.commit()
                    continue

                _attempt_charge(db, record, billing_method)

        except Exception as exc:
            logger.error("Error processing billing for merchant %s: %s", merchant_id, exc, exc_info=True)
            failed_count += 1

    logger.info(
        "generate_platform_billing done: created=%d skipped=%d failed=%d",
        created_count, skipped_count, failed_count,
    )
    return {"created": created_count, "skipped": skipped_count, "failed": failed_count}


def _attempt_charge(db, record, billing_method):
    """Attempt to charge the merchant's billing method. Updates record status in-place."""
    from src.apps.pricing_control.models.billing_record import BillingStatus

    if record.total_due_cents <= 0:
        record.status = BillingStatus.SKIPPED
        db.commit()
        return

    try:
        record.status = BillingStatus.PROCESSING
        db.commit()

        # Platform charge via TSYS — uses HW_TSYS_* env vars
        # Deferred to avoid circular imports; returns charge_txn_id on success
        charge_txn_id = _submit_platform_charge(
            token=billing_method.tsep_token,
            amount_cents=record.total_due_cents,
            merchant_id=record.merchant_id,
            billing_record_id=record.id,
        )

        record.status = BillingStatus.PAID
        record.charge_txn_id = charge_txn_id
        db.commit()

    except Exception as charge_exc:
        logger.error("Charge failed for merchant %s billing_record %s: %s", record.merchant_id, record.id, charge_exc)
        record.retry_count = (record.retry_count or 0) + 1
        record.last_error = str(charge_exc)[:500]
        if record.retry_count >= 3:
            record.status = BillingStatus.FAILED
        else:
            record.status = BillingStatus.RETRYING
        db.commit()


def _submit_platform_charge(token: str, amount_cents: int, merchant_id: int, billing_record_id: int) -> str:
    """
    Submit charge via HW platform TSYS account.
    Uses HW_TSYS_* environment variables (platform account, not merchant credentials).
    Returns provider transaction ID on success. Raises on failure.
    """
    from src.core.config import settings

    hw_tsys_key = getattr(settings, "HW_TSYS_API_KEY", None)
    if not hw_tsys_key:
        raise RuntimeError("HW_TSYS_API_KEY not configured — platform billing charge skipped")

    # NOTE: Actual TSYS charge submission is deferred to when HW_TSYS platform
    # credentials are wired in (PRD-HWMP-004). For now, log and return a placeholder.
    logger.info(
        "PLATFORM_CHARGE merchant=%s amount_cents=%d billing_record=%d [HW_TSYS integration pending]",
        merchant_id, amount_cents, billing_record_id,
    )
    # TODO: Implement full TSYS charge submission when HW_TSYS credentials are available.
    # charge_result = tsys_client.charge(token=token, amount=amount_cents, ...)
    # return charge_result.transaction_id
    raise NotImplementedError("HW_TSYS platform charge not yet implemented — pending PRD-HWMP-004")
