"""
Celery task: catch-all for PENDING transactions that missed their Celery ETA.

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

Queries PENDING transactions joined to invoices where the invoice due_date
is today or earlier.  For each, re-queues process_scheduled_payment immediately
(no ETA).  Skips transactions processed within the last 23 hours.

Returns: {"found": N, "dispatched": N, "skipped_recent": N}
"""
import logging
from datetime import datetime, timedelta, timezone

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__)


@celery_app.task(
    bind=True,
    name="transactions.process_due_transactions",
    max_retries=3,
    default_retry_delay=300,
)
def process_due_transactions(self) -> dict:
    """
    Daily catch-all at 08:00 UTC.

    Finds PENDING transactions whose linked invoice is due today or earlier
    and re-fires process_scheduled_payment for each.
    """
    from sqlalchemy import select, join
    from src.apps.transactions.models.transactions import Transactions, transactions_invoices_map
    from src.apps.invoices.models.invoice import Invoice
    from src.core.utils.enums import TransactionStatusTypes, TransactionTypes

    now = datetime.now(timezone.utc)
    today = now.date()

    found = 0
    dispatched = 0
    skipped_recent = 0

    try:
        with SessionCelery() as db:
            # Find PENDING transactions joined to invoices where due_date <= today.
            # process_scheduled_payment resolves the provider from txn.merchant_id — cross-merchant
            # risk is mitigated by the FK constraints on payment_method_id and the provider lookup.
            # merchant_id is logged here for audit traceability.
            stmt = (
                select(Transactions.id, Transactions.payment_method_id, Transactions.txn_metadata, Transactions.merchant_id)
                .join(
                    transactions_invoices_map,
                    transactions_invoices_map.c.transaction_id == Transactions.id,
                )
                .join(
                    Invoice,
                    Invoice.id == transactions_invoices_map.c.invoice_id,
                )
                .where(
                    Transactions.txn_status == TransactionStatusTypes.PENDING,
                    Invoice.due_date <= now,
                    Invoice.deleted_at.is_(None),
                    Transactions.transaction_type.in_([
                        TransactionTypes.SUBSCRIPTION,
                        TransactionTypes.PAYMENT_TERMINAL,
                    ]),
                )
                .distinct()
            )
            rows = db.execute(stmt).all()
            txn_ids = [(r[0], r[1], r[2], r[3]) for r in rows]

    except Exception as exc:
        logger.error("process_due_transactions: fatal fetch error: %s", exc, exc_info=True)
        raise self.retry(exc=exc, countdown=300)

    cutoff = now - timedelta(hours=23)

    for txn_id, pm_id, txn_metadata, merchant_id in txn_ids:
        found += 1
        try:
            # Check 23-hour skip guard
            last_attempt = (txn_metadata or {}).get("last_catch_all_attempt_at")
            if last_attempt:
                try:
                    last_dt = datetime.fromisoformat(last_attempt)
                    if last_dt.tzinfo is None:
                        last_dt = last_dt.replace(tzinfo=timezone.utc)
                    if last_dt >= cutoff:
                        skipped_recent += 1
                        continue
                except (ValueError, TypeError):
                    pass

            # Update metadata
            with SessionCelery() as db:
                from sqlalchemy import select as _sel
                from src.apps.transactions.models.transactions import Transactions as _T
                txn = db.execute(_sel(_T).where(_T.id == txn_id)).scalar_one_or_none()
                if txn is None:
                    continue
                txn.txn_metadata = {**(txn.txn_metadata or {}), "last_catch_all_attempt_at": now.isoformat()}
                db.commit()

            # Re-queue payment
            if pm_id:
                from src.worker.hpp_tasks import process_scheduled_payment
                process_scheduled_payment.apply_async(args=[txn_id, pm_id])
                dispatched += 1
                logger.info("process_due_transactions: re-queued txn_id=%s merchant_id=%s", txn_id, merchant_id)
            else:
                logger.warning(
                    "process_due_transactions: txn_id=%s has no payment_method_id, skipping",
                    txn_id,
                )

        except Exception as txn_err:
            logger.error(
                "process_due_transactions: error for txn_id=%s: %s", txn_id, txn_err, exc_info=True
            )

    logger.info(
        "process_due_transactions: found=%d dispatched=%d skipped_recent=%d",
        found, dispatched, skipped_recent,
    )
    return {"found": found, "dispatched": dispatched, "skipped_recent": skipped_recent}
