# NOTE: These tasks process invoices across ALL merchants intentionally.
# Per-merchant isolation is enforced by process_scheduled_payment (which validates
# the payment method belongs to the transaction's merchant) and all write paths.

"""
Celery task: automatic dunning retry dispatcher.

Task name: transactions.dunning_retry_dispatcher
Schedule:  daily at 08:30 UTC (Celery Beat)
Queue:     scheduler

Three sequential parts per daily invocation:

  Part A — Schedule newly failed invoices:
    FAILED invoices with retry_count < max_retries and next_retry_at IS NULL
    → compute next_retry_at, emit invoice.dunning_retry_scheduled event

  Part B — Execute due retries:
    FAILED invoices with next_retry_at <= now and retry_count < max_retries
    → find or create a PENDING Transaction, fire process_scheduled_payment,
      increment retry_count, clear next_retry_at

  Part C — Exhaust dunning:
    FAILED invoices with retry_count >= max_retries and retry_exhausted_at IS NULL
    → set retry_exhausted_at, transition subscription to DUNNING_EXHAUSTED,
      emit invoice.dunning_exhausted event

NOTE: Same-day invoices (billing_date == created_at.date()) are skipped in
Part B and logged as WARNING — see PRD-010 §3 for rationale.

Returns: {"scheduled": N, "retried": N, "exhausted": N, "errors": N}
"""
import asyncio
import logging
from datetime import datetime, timedelta, timezone
from typing import Optional

from celery.utils.log import get_task_logger

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

logger = get_task_logger(__name__)

# Retry schedule: retry_count 0 → T+3 days, retry_count 1 → T+7 days.
# For retry_count >= len(DUNNING_RETRY_DAYS), wrap around with modulo.
DUNNING_RETRY_DAYS = [3, 7]


@celery_app.task(
    bind=True,
    name="transactions.dunning_retry_dispatcher",
    max_retries=3,
    default_retry_delay=300,
)
def dunning_retry_dispatcher(self) -> dict:
    """Daily dunning retry orchestrator at 08:30 UTC."""
    now = datetime.now(timezone.utc)

    total_scheduled = 0
    total_retried = 0
    total_exhausted = 0
    total_errors = 0

    s, e = _part_a_schedule_failed(now)
    total_scheduled += s
    total_errors += e

    r, e = _part_b_execute_retries(now)
    total_retried += r
    total_errors += e

    ex, e = _part_c_exhaust_dunning(now)
    total_exhausted += ex
    total_errors += e

    logger.info(
        "dunning_retry_dispatcher: scheduled=%d retried=%d exhausted=%d errors=%d",
        total_scheduled, total_retried, total_exhausted, total_errors,
    )
    return {
        "scheduled": total_scheduled,
        "retried": total_retried,
        "exhausted": total_exhausted,
        "errors": total_errors,
    }


def _part_a_schedule_failed(now: datetime):
    """Schedule next_retry_at for newly FAILED invoices."""
    from sqlalchemy import select
    from src.apps.invoices.models.invoice import Invoice
    from src.core.utils.enums import InvoiceStatusTypes, InvoiceActivityTypes
    from src.apps.invoices import crud as invoice_crud

    scheduled = 0
    errors = 0
    event_payloads = []

    try:
        with SessionCelery() as db:
            stmt = select(Invoice).where(
                Invoice.status == InvoiceStatusTypes.FAILED,
                Invoice.retry_count < Invoice.max_retries,
                Invoice.next_retry_at.is_(None),
                Invoice.deleted_at.is_(None),
            ).limit(500)
            invoices = db.execute(stmt).scalars().all()

            if len(invoices) == 500:
                logger.warning(
                    "_part_a_schedule_failed: hit 500-row batch limit — additional invoices will be processed next run"
                )

            for invoice in invoices:
                try:
                    logger.debug(
                        "_part_a: scheduling retry for invoice %s merchant_id=%s",
                        invoice.id, invoice.merchant_id
                    )
                    days_offset = DUNNING_RETRY_DAYS[invoice.retry_count % len(DUNNING_RETRY_DAYS)]
                    base_time = invoice.last_retry_at or invoice.updated_at or now
                    if base_time and base_time.tzinfo is None:
                        base_time = base_time.replace(tzinfo=timezone.utc)
                    next_retry = base_time + timedelta(days=days_offset)
                    invoice.next_retry_at = next_retry
                    invoice.updated_at = now

                    invoice_crud.write_activity(
                        db=db,
                        invoice_id=invoice.id,
                        activity_type=InvoiceActivityTypes.INVOICE_UPDATED,
                        description=f"Dunning retry scheduled (attempt {invoice.retry_count + 1})",
                        actor_type="system",
                        metadata={
                            "retry_number": invoice.retry_count + 1,
                            "retry_at": next_retry.isoformat(),
                        },
                    )

                    event_payloads.append({
                        "invoice_id": invoice.id,
                        "invoice_literal": invoice.invoice_literal,
                        "merchant_id": invoice.merchant_id,
                        "customer_id": invoice.customer_id,
                        "retry_number": invoice.retry_count + 1,
                        "retry_at": next_retry.isoformat(),
                    })
                    scheduled += 1

                except Exception as inv_err:
                    logger.error("_part_a_schedule_failed: invoice %s: %s", invoice.id, inv_err)
                    errors += 1

    except Exception as exc:
        logger.error("_part_a_schedule_failed: fatal error: %s", exc, exc_info=True)
        errors += 1
        return scheduled, errors

    if event_payloads:
        _emit_events("invoice.dunning_retry_scheduled", event_payloads)

    return scheduled, errors


def _part_b_execute_retries(now: datetime):
    """Execute dunning retries that are now due."""
    from sqlalchemy import select
    from src.apps.invoices.models.invoice import Invoice
    from src.apps.transactions.models.transactions import Transactions, transactions_invoices_map
    from src.core.utils.enums import (
        InvoiceStatusTypes, InvoiceActivityTypes, TransactionStatusTypes,
        TransactionTypes, TransactionCategories
    )
    from src.apps.invoices import crud as invoice_crud

    retried = 0
    errors = 0

    try:
        with SessionCelery() as db:
            stmt = select(Invoice).where(
                Invoice.status == InvoiceStatusTypes.FAILED,
                Invoice.next_retry_at <= now,
                Invoice.retry_count < Invoice.max_retries,
                Invoice.deleted_at.is_(None),
            )
            invoices = db.execute(stmt).scalars().all()
            invoice_ids = [inv.id for inv in invoices]

    except Exception as exc:
        logger.error("_part_b_execute_retries: fatal fetch error: %s", exc, exc_info=True)
        return 0, 1

    for invoice_id in invoice_ids:
        try:
            with SessionCelery() as db:
                from sqlalchemy import select as _sel
                inv_stmt = _sel(Invoice).where(Invoice.id == invoice_id)
                invoice = db.execute(inv_stmt).scalar_one_or_none()
                if invoice is None:
                    continue
                # Re-validate status — another process may have already handled this invoice
                if invoice.status != InvoiceStatusTypes.FAILED:
                    logger.info(
                        "_part_b_execute_retries: invoice %s status changed to %s — skipping",
                        invoice_id, invoice.status
                    )
                    continue

                # Skip same-day invoices
                billing_date = getattr(invoice, "billing_date", None)
                created_date = invoice.created_at.date() if invoice.created_at else None
                billing_day = billing_date.date() if billing_date else None
                if billing_day and created_date and billing_day == created_date:
                    logger.warning(
                        "dunning_retry_dispatcher: skipping same-day invoice %s "
                        "(billing_date==created_at.date())",
                        invoice.id,
                    )
                    continue

                # Find most recent FAILED or PENDING transaction
                txn_stmt = (
                    _sel(Transactions)
                    .join(
                        transactions_invoices_map,
                        transactions_invoices_map.c.transaction_id == Transactions.id,
                    )
                    .where(
                        transactions_invoices_map.c.invoice_id == invoice_id,
                        Transactions.txn_status.in_([
                            TransactionStatusTypes.FAILED,
                            TransactionStatusTypes.PENDING,
                        ]),
                    )
                    .order_by(Transactions.id.desc())
                )
                existing_txn = db.execute(txn_stmt).scalars().first()

                if existing_txn:
                    txn = existing_txn
                    pm_id = txn.payment_method_id
                else:
                    # Create new PENDING transaction
                    import uuid as _uuid
                    from src.apps.transactions.services import generate_txn_literal
                    txn = Transactions(
                        txn_id=f"dun_{invoice.invoice_id}_{_uuid.uuid4().hex[:8]}",
                        txn_literal=generate_txn_literal(db),
                        txn_amount=invoice.amount,
                        txn_status=TransactionStatusTypes.PENDING,
                        payment_request_id=invoice.payment_request_id,
                        merchant_id=invoice.merchant_id,
                        customer_id=invoice.customer_id,
                        transaction_type=TransactionTypes.PAYMENT_TERMINAL,
                        category=TransactionCategories.CHARGE,
                        txn_metadata={"dunning_retry": True},
                    )
                    db.add(txn)
                    db.flush()
                    db.execute(
                        transactions_invoices_map.insert().values(
                            transaction_id=txn.id,
                            invoice_id=invoice_id,
                        )
                    )
                    db.flush()
                    pm_id = None

                # Guard: only proceed if we have a payment method to charge.
                # Check BEFORE committing so we never persist a retry_count increment
                # without an actual charge attempt (avoids crash-window undo complexity).
                if pm_id is None:
                    logger.warning(
                        "dunning_retry_dispatcher: invoice %s has no payment_method_id — "
                        "skipping retry, state unchanged",
                        invoice_id,
                    )
                    continue

                # Update invoice dunning state
                invoice.retry_count = (invoice.retry_count or 0) + 1
                invoice.last_retry_at = now
                invoice.next_retry_at = None
                invoice.status = InvoiceStatusTypes.PENDING
                invoice.updated_at = now

                # Update subscription dunning counters if applicable
                if invoice.subscription_id:
                    from src.apps.subscriptions.models.subscription import Subscription
                    sub = db.execute(
                        _sel(Subscription).where(Subscription.id == invoice.subscription_id)
                    ).scalar_one_or_none()
                    if sub:
                        sub.dunning_retry_count = (getattr(sub, "dunning_retry_count", 0) or 0) + 1
                        sub.dunning_last_retry_at = now

                invoice_crud.write_activity(
                    db=db,
                    invoice_id=invoice_id,
                    activity_type=InvoiceActivityTypes.INVOICE_UPDATED,
                    description=f"Dunning retry attempt {invoice.retry_count}",
                    actor_type="system",
                    metadata={"retry_count": invoice.retry_count},
                )
                db.commit()

                # Revoke any existing scheduled Celery ETA task to prevent double-charge
                existing_task_id = (txn.txn_metadata or {}).get("celery_task_id")
                if existing_task_id:
                    try:
                        from src.worker.celery_app import celery_app as _celery_app
                        _celery_app.control.revoke(existing_task_id, terminate=False)
                        logger.info(
                            "_part_b_execute_retries: revoked old ETA task %s for txn %s",
                            existing_task_id, txn.id,
                        )
                    except Exception as revoke_err:
                        logger.warning(
                            "_part_b_execute_retries: failed to revoke task %s: %s",
                            existing_task_id, revoke_err
                        )

                from src.worker.hpp_tasks import process_scheduled_payment
                process_scheduled_payment.apply_async(args=[txn.id, pm_id])
                retried += 1
                logger.info(
                    "dunning_retry_dispatcher: queued retry for invoice %s (attempt %d)",
                    invoice_id, invoice.retry_count
                )

        except Exception as inv_err:
            logger.error(
                "_part_b_execute_retries: invoice %s: %s", invoice_id, inv_err, exc_info=True
            )
            errors += 1

    return retried, errors


def _part_c_exhaust_dunning(now: datetime):
    """Mark dunning-exhausted invoices and transition subscriptions."""
    from sqlalchemy import select
    from src.apps.invoices.models.invoice import Invoice
    from src.core.utils.enums import InvoiceStatusTypes, InvoiceActivityTypes
    from src.apps.invoices import crud as invoice_crud

    exhausted = 0
    errors = 0
    event_payloads = []

    try:
        with SessionCelery() as db:
            stmt = select(Invoice).where(
                Invoice.status == InvoiceStatusTypes.FAILED,
                Invoice.retry_count >= Invoice.max_retries,
                Invoice.retry_exhausted_at.is_(None),
                Invoice.deleted_at.is_(None),
            ).limit(500)
            invoices = db.execute(stmt).scalars().all()

            if len(invoices) == 500:
                logger.warning(
                    "_part_c_exhaust_dunning: hit 500-row batch limit — additional invoices will be processed next run"
                )

            for invoice in invoices:
                try:
                    logger.debug(
                        "_part_c: exhausting invoice %s merchant_id=%s",
                        invoice.id, invoice.merchant_id
                    )
                    invoice.retry_exhausted_at = now
                    invoice.updated_at = now

                    # Transition subscription to DUNNING_EXHAUSTED if applicable
                    if invoice.subscription_id:
                        from src.apps.subscriptions.models.subscription import Subscription
                        from src.apps.subscriptions.enums import SubscriptionStatus
                        from sqlalchemy import select as _sel
                        sub = db.execute(
                            _sel(Subscription).where(Subscription.id == invoice.subscription_id)
                        ).scalar_one_or_none()
                        if sub:
                            past_due_val = getattr(SubscriptionStatus, "PAST_DUE", None)
                            dunning_exhausted_val = getattr(SubscriptionStatus, "DUNNING_EXHAUSTED", None)
                            if (
                                past_due_val is not None
                                and dunning_exhausted_val is not None
                                and sub.status == past_due_val
                            ):
                                sub.status = dunning_exhausted_val
                                sub.updated_at = now
                                logger.info(
                                    "dunning_retry_dispatcher: subscription %s → DUNNING_EXHAUSTED",
                                    sub.subscription_literal,
                                )

                    invoice_crud.write_activity(
                        db=db,
                        invoice_id=invoice.id,
                        activity_type=InvoiceActivityTypes.INVOICE_UPDATED,
                        description="Dunning exhausted: all retry attempts failed",
                        actor_type="system",
                        metadata={"retry_count": invoice.retry_count},
                    )

                    event_payloads.append({
                        "invoice_id": invoice.id,
                        "invoice_literal": invoice.invoice_literal,
                        "merchant_id": invoice.merchant_id,
                        "customer_id": invoice.customer_id,
                        "amount": float(invoice.amount or 0),
                        "retry_count": invoice.retry_count,
                        "exhausted_at": now.isoformat(),
                    })
                    exhausted += 1

                except Exception as inv_err:
                    logger.error("_part_c_exhaust_dunning: invoice %s: %s", invoice.id, inv_err)
                    errors += 1

    except Exception as exc:
        logger.error("_part_c_exhaust_dunning: fatal error: %s", exc, exc_info=True)
        errors += 1
        return exhausted, errors

    if event_payloads:
        _emit_events("invoice.dunning_exhausted", event_payloads)

    return exhausted, errors


def _emit_events(event_type: str, payloads: list) -> None:
    """Emit one Kafka event per payload entry."""
    from src.events.base import BaseEvent
    from src.events.dispatcher import EventDispatcher

    async def _dispatch_all():
        for p in payloads:
            await EventDispatcher.dispatch(BaseEvent(event_type=event_type, data=p))

    try:
        loop = asyncio.new_event_loop()
        try:
            loop.run_until_complete(_dispatch_all())
        finally:
            loop.close()
    except Exception as exc:
        logger.error("_emit_events(%s): failed: %s", event_type, exc)
