"""
Daily Celery task — process payment requests whose billing/due date is today.

Task name:  payment_requests.process_due_payment_requests
Schedule:   daily at 08:00 UTC (Celery Beat)

The task is fully idempotent — running it twice for the same date is safe
because each processed PR is flipped out of PENDING status immediately.

Handles:
  - ONE_TIME / no-frequency payment requests: charges the full PR amount.
  - SPLIT payment requests: charges each installment (split_payment_requests row)
    whose billing_date / due_date falls on today.
  - RECURRING payment requests: skipped — handled by the subscription tasks.

Each item is processed in its own DB session so a single failure never
rolls back unrelated successful charges committed earlier in the batch.
"""

import asyncio
import logging
from datetime import datetime, timezone
from typing import Optional

from celery.utils.log import get_task_logger

from src.worker.celery_app import celery_app

logger = get_task_logger(__name__)


# ---------------------------------------------------------------------------
# Main Beat task
# ---------------------------------------------------------------------------

@celery_app.task(
    bind=True,
    name="payment_requests.process_due_payment_requests",
    max_retries=0,  # batch task — individual item errors are logged, not retried globally
)
def process_due_payment_requests(self, **kwargs) -> dict:
    """
    Find all PENDING payment requests / split installments whose
    billing_date (or due_date when billing_date is NULL) matches today
    and execute the charge for each via the active payment provider.
    """
    from sqlalchemy import select, or_, and_, func
    from src.core.database import SessionCelery
    from src.apps.payment_requests.models.payment_request import PaymentRequest
    from src.apps.payment_requests.models.split_payment_requests import SplitPaymentRequests
    from src.apps.payment_requests.enums import PaymentRequestStatusTypes, PaymentFrequencies

    today = datetime.now(timezone.utc).date()
    logger.info("process_due_payment_requests: running for date=%s", today)

    counts = {"processed": 0, "succeeded": 0, "failed": 0}

    # ── Collect due IDs in a lightweight read-only pass ─────────────────────
    # IDs are stored before closing the session so each charge can open its
    # own independent session without holding a long-lived connection.
    with SessionCelery() as db:
        # One-time (and frequency=None) payment requests
        due_pr_ids = db.execute(
            select(PaymentRequest.id).where(
                PaymentRequest.status == PaymentRequestStatusTypes.PENDING,
                PaymentRequest.deleted_at.is_(None),
                # Skip RECURRING — subscription tasks own those
                or_(
                    PaymentRequest.payment_frequency == PaymentFrequencies.ONE_TIME,
                    PaymentRequest.payment_frequency.is_(None),
                ),
                # billing_date = today, or (billing_date NULL and due_date = today)
                or_(
                    func.date(PaymentRequest.billing_date) == today,
                    and_(
                        PaymentRequest.billing_date.is_(None),
                        func.date(PaymentRequest.due_date) == today,
                    ),
                ),
            )
        ).scalars().all()

        # Split installments due today (parent PR must still be PENDING)
        due_split_ids = db.execute(
            select(SplitPaymentRequests.id).join(
                PaymentRequest,
                SplitPaymentRequests.payment_request_id == PaymentRequest.id,
            ).where(
                PaymentRequest.status == PaymentRequestStatusTypes.PENDING,
                PaymentRequest.deleted_at.is_(None),
                SplitPaymentRequests.paid_date.is_(None),
                or_(
                    func.date(SplitPaymentRequests.billing_date) == today,
                    and_(
                        SplitPaymentRequests.billing_date.is_(None),
                        func.date(SplitPaymentRequests.due_date) == today,
                    ),
                ),
            )
        ).scalars().all()

    logger.info(
        "process_due_payment_requests: %d one-time PRs, %d split installments due today",
        len(due_pr_ids),
        len(due_split_ids),
    )

    # ── Process one-time payment requests ───────────────────────────────────
    for pr_id in due_pr_ids:
        counts["processed"] += 1
        try:
            success = _charge_payment_request(pr_id)
            counts["succeeded" if success else "failed"] += 1
        except Exception as exc:
            logger.error(
                "process_due_payment_requests: unhandled error on PR id=%s: %s",
                pr_id,
                exc,
                exc_info=True,
            )
            counts["failed"] += 1

    # ── Process split installments ───────────────────────────────────────────
    for split_id in due_split_ids:
        counts["processed"] += 1
        try:
            success = _charge_split_installment(split_id)
            counts["succeeded" if success else "failed"] += 1
        except Exception as exc:
            logger.error(
                "process_due_payment_requests: unhandled error on split id=%s: %s",
                split_id,
                exc,
                exc_info=True,
            )
            counts["failed"] += 1

    logger.info(
        "process_due_payment_requests: done — processed=%d succeeded=%d failed=%d",
        counts["processed"],
        counts["succeeded"],
        counts["failed"],
    )
    return counts


# ---------------------------------------------------------------------------
# Shared helpers
# ---------------------------------------------------------------------------

def _run_provider_charge(merchant, db, pm, amount_dollars: float, currency: str,
                          idempotency_key: str, customer_id: int):
    """
    Resolve the active payment provider for *merchant* and submit a charge.

    Runs the async provider call inside a fresh event loop so this helper
    can be called from a synchronous Celery task.

    Returns:
        (provider_config, charge_result_obj)  — the raw provider objects.

    Raises:
        Exception if the provider cannot be resolved or the HTTP call fails.
    """
    from src.core.providers.factory import get_provider_for_merchant

    async def _run():
        # Ensure provider implementations are registered. Celery workers do not
        # run FastAPI startup hooks, so the @ProviderRegistry.register decorators
        # must be triggered manually before the first provider lookup.
        from src.core.providers.loader import load_all_providers
        load_all_providers()

        active_provider, provider_config = await get_provider_for_merchant(merchant, db)
        payment_method_type = pm.method or "card"

        if payment_method_type == "ach" and hasattr(active_provider, "submit_ach_charge"):
            ach = pm.ach_details
            routing = ach.routing_number if ach else ""
            account_type = ach.account_type if ach else "CHECKING"
            name = ach.account_name if ach else ""
            parts = name.split(" ", 1)
            result = await active_provider.submit_ach_charge(
                config=provider_config,
                routing_number=routing,
                account_number=pm.reference_id,
                account_type=account_type,
                first_name=parts[0] if parts else "",
                last_name=parts[1] if len(parts) > 1 else "",
                amount=amount_dollars,
            )
        else:
            result = await active_provider.submit_charge(
                config=provider_config,
                amount=amount_dollars,
                currency=currency,
                payment_method_token=pm.reference_id,
                payment_method_type=payment_method_type,
                capture=True,
                idempotency_key=idempotency_key,
                metadata={
                    "merchant_id": merchant.id,
                    "customer_id": customer_id,
                    "scheduled": True,
                },
            )
        return provider_config, result

    loop = asyncio.new_event_loop()
    try:
        return loop.run_until_complete(_run())
    finally:
        loop.close()


def _persist_provider_transaction(db, transaction, provider_config, charge_result: dict) -> None:
    """
    Create a ProviderTransaction row for a scheduled charge.
    Mirrors the VT / Invoice charge path in payment_requests/services.py.
    """
    try:
        from src.apps.transactions.models.provider_transactions import ProviderTransaction

        raw = charge_result.get("raw") or {}
        sale_resp = (
            raw.get("SaleResponse")
            or raw.get("PreAuthResponse")
            or raw.get("AchResponse")
            or {}
        )
        pt = ProviderTransaction(
            transaction_id=transaction.id,
            provider_slug=provider_config.provider_slug,
            provider_txn_id=(
                sale_resp.get("transactionID") or charge_result.get("transaction_id")
            ),
            host_reference_number=sale_resp.get("hostReferenceNumber"),
            response_code=(
                sale_resp.get("approvalCode") or sale_resp.get("responseCode")
            ),
            authorization_code=(
                sale_resp.get("authorizationCode") or sale_resp.get("authCode")
            ),
            card_type=sale_resp.get("cardType"),
            masked_card_number=sale_resp.get("maskedCardNumber"),
            avs_response_code=sale_resp.get("AVSResponseCode"),
            cvv_response_code=sale_resp.get("CVVResponseCode"),
            amount=charge_result.get("amount"),
            currency=charge_result.get("currency"),
            status=charge_result.get("status"),
            raw_response=raw,
        )
        db.add(pt)
        db.flush()
    except Exception as exc:
        logger.warning(
            "_persist_provider_transaction: could not create ProviderTransaction: %s", exc
        )


# ---------------------------------------------------------------------------
# Per-item charge helpers (each opens its own DB session)
# ---------------------------------------------------------------------------

def _charge_payment_request(pr_id: int) -> bool:
    """
    Load a single PENDING payment request, submit the charge via the active
    provider, create a Transaction record and mark the PR as PAID / FAILED.

    Returns True on a successful charge.
    """
    from sqlalchemy import select
    from sqlalchemy.orm import selectinload
    from src.core.database import SessionCelery
    from src.core.utils.enums import TransactionStatusTypes, TransactionCategories, TransactionTypes
    from src.apps.base.utils.functions import generate_secure_id
    from src.apps.transactions.services import create_transaction
    from src.apps.payment_requests.models.payment_request import PaymentRequest
    from src.apps.payment_requests.models.payment_request_customer import PaymentRequestCustomer
    from src.apps.payment_requests.enums import PaymentRequestStatusTypes
    from src.apps.merchants.models.merchant import Merchant

    with SessionCelery() as db:
        pr = db.execute(
            select(PaymentRequest)
            .where(PaymentRequest.id == pr_id)
            .options(selectinload(PaymentRequest.payment_methods))
        ).scalar_one_or_none()

        if not pr:
            logger.warning("_charge_payment_request: PR id=%s not found — skipping", pr_id)
            return False

        # Guard: re-check status (another process may have handled it between ID collection and now)
        if pr.status != PaymentRequestStatusTypes.PENDING:
            logger.info(
                "_charge_payment_request: PR %s already in status=%s — skipping",
                pr.payment_request_id,
                pr.status,
            )
            return True  # not an error

        # Resolve customer
        pr_customer = db.execute(
            select(PaymentRequestCustomer).where(
                PaymentRequestCustomer.payment_request_id == pr.id
            )
        ).scalar_one_or_none()
        if not pr_customer:
            logger.error(
                "_charge_payment_request: PR %s has no customer record — skipping",
                pr.payment_request_id,
            )
            return False

        # Resolve payment method
        if not pr.payment_methods:
            logger.error(
                "_charge_payment_request: PR %s has no payment methods — skipping",
                pr.payment_request_id,
            )
            return False
        pm = pr.payment_methods[0]

        # Resolve provider token
        try:
            provider_token = pm.reference_id
        except Exception:
            provider_token = None
        if not provider_token:
            logger.error(
                "_charge_payment_request: PR %s — PM %s has no provider token — skipping",
                pr.payment_request_id,
                pm.id,
            )
            return False

        # Resolve merchant
        merchant = db.execute(
            select(Merchant).where(Merchant.id == pr.merchant_id)
        ).scalar_one_or_none()
        if not merchant:
            logger.error(
                "_charge_payment_request: merchant id=%s not found — skipping",
                pr.merchant_id,
            )
            return False

        amount_dollars = float(pr.amount) / 100  # PR amounts are stored in cents
        currency = pr.currency.lower() if pr.currency else "usd"

        # Submit charge
        try:
            provider_config, charge_result_obj = _run_provider_charge(
                merchant=merchant,
                db=db,
                pm=pm,
                amount_dollars=amount_dollars,
                currency=currency,
                idempotency_key=pr.payment_request_id,
                customer_id=pr_customer.customer_id,
            )
        except Exception as charge_exc:
            logger.error(
                "_charge_payment_request: provider error for PR %s: %s",
                pr.payment_request_id,
                charge_exc,
                exc_info=True,
            )
            pr.status = PaymentRequestStatusTypes.FAILED
            return False

        charge_result = {
            "status": charge_result_obj.status,
            "transaction_id": charge_result_obj.transaction_id,
            "amount": charge_result_obj.amount,
            "currency": charge_result_obj.currency,
            "raw": charge_result_obj.raw_response,
        }

        if charge_result.get("status") == "succeeded":
            pr.status = PaymentRequestStatusTypes.PAID

            txn_id = charge_result.get("transaction_id") or generate_secure_id(
                prepend="txn", length=20
            )
            transaction = create_transaction(
                db,
                txn_id=txn_id,
                txn_amount=float(pr.amount),  # cents, consistent with services.py
                currency=currency,
                txn_status=TransactionStatusTypes.PAID,
                txn_type=pm.method or "card",
                txn_source="scheduled",
                category=TransactionCategories.CHARGE,
                transaction_type=TransactionTypes.PAYMENT_TERMINAL,
                txn_metadata=charge_result,
                payment_request_id=pr.id,
                merchant_id=merchant.id,
                customer_id=pr_customer.customer_id,
                payment_method_id=pm.id,
            )
            _persist_provider_transaction(db, transaction, provider_config, charge_result)
            # SessionCelery commits on context-manager exit
            logger.info(
                "_charge_payment_request: PR %s → PAID (txn_id=%s)",
                pr.payment_request_id,
                txn_id,
            )
            return True
        else:
            pr.status = PaymentRequestStatusTypes.FAILED
            _fail_txn_id = charge_result.get("transaction_id") or generate_secure_id(prepend="txn", length=20)
            fail_transaction = create_transaction(
                db,
                txn_id=_fail_txn_id,
                txn_amount=float(pr.amount),
                currency=currency,
                txn_status=TransactionStatusTypes.FAILED,
                txn_type=pm.method or "card",
                txn_source="scheduled",
                category=TransactionCategories.CHARGE,
                transaction_type=TransactionTypes.PAYMENT_TERMINAL,
                txn_metadata=charge_result,
                payment_request_id=pr.id,
                merchant_id=merchant.id,
                customer_id=pr_customer.customer_id,
                payment_method_id=pm.id,
            )
            _persist_provider_transaction(db, fail_transaction, provider_config, charge_result)
            logger.warning(
                "_charge_payment_request: PR %s charge FAILED (provider status=%s)",
                pr.payment_request_id,
                charge_result.get("status"),
            )
            return False


def _charge_split_installment(split_id: int) -> bool:
    """
    Load a single unpaid split installment, submit the charge, create a
    Transaction record, mark the split as paid, and update the parent PR
    status to PAID if all installments are now paid.

    Returns True on a successful charge.
    """
    from sqlalchemy import select
    from sqlalchemy.orm import selectinload
    from src.core.database import SessionCelery
    from src.core.utils.enums import TransactionStatusTypes, TransactionCategories, TransactionTypes
    from src.apps.base.utils.functions import generate_secure_id
    from src.apps.transactions.services import create_transaction
    from src.apps.payment_requests.models.payment_request import PaymentRequest
    from src.apps.payment_requests.models.split_payment_requests import SplitPaymentRequests
    from src.apps.payment_requests.models.payment_request_customer import PaymentRequestCustomer
    from src.apps.payment_requests.enums import PaymentRequestStatusTypes, SplitPaymentTypes
    from src.apps.merchants.models.merchant import Merchant

    with SessionCelery() as db:
        split = db.execute(
            select(SplitPaymentRequests).where(SplitPaymentRequests.id == split_id)
        ).scalar_one_or_none()

        if not split:
            logger.warning("_charge_split_installment: split id=%s not found — skipping", split_id)
            return False

        # Guard: skip if already paid
        if split.paid_date is not None:
            logger.info(
                "_charge_split_installment: split %s already paid — skipping", split_id
            )
            return True

        # Load parent PR
        pr = db.execute(
            select(PaymentRequest)
            .where(PaymentRequest.id == split.payment_request_id)
            .options(selectinload(PaymentRequest.payment_methods))
        ).scalar_one_or_none()

        if not pr or pr.status != PaymentRequestStatusTypes.PENDING:
            logger.warning(
                "_charge_split_installment: split %s — parent PR not found or not PENDING",
                split_id,
            )
            return False

        # Resolve customer
        pr_customer = db.execute(
            select(PaymentRequestCustomer).where(
                PaymentRequestCustomer.payment_request_id == pr.id
            )
        ).scalar_one_or_none()
        if not pr_customer:
            logger.error(
                "_charge_split_installment: split %s — PR %s has no customer",
                split_id,
                pr.payment_request_id,
            )
            return False

        # Resolve payment method
        if not pr.payment_methods:
            logger.error(
                "_charge_split_installment: split %s — PR %s has no payment methods",
                split_id,
                pr.payment_request_id,
            )
            return False
        pm = pr.payment_methods[0]

        try:
            provider_token = pm.reference_id
        except Exception:
            provider_token = None
        if not provider_token:
            logger.error(
                "_charge_split_installment: split %s — PM %s has no provider token",
                split_id,
                pm.id,
            )
            return False

        # Compute installment amount (keep in cents for txn_amount; dollars for provider)
        if split.split_type == SplitPaymentTypes.PERCENTAGE:
            txn_amount_cents = round((split.split_value / 100.0) * pr.amount)
        else:  # AMOUNT — split_value is already in cents
            txn_amount_cents = split.split_value
        amount_dollars = txn_amount_cents / 100.0

        if amount_dollars <= 0:
            logger.error(
                "_charge_split_installment: split %s — invalid amount %.2f — skipping",
                split_id,
                amount_dollars,
            )
            return False

        # Resolve merchant
        merchant = db.execute(
            select(Merchant).where(Merchant.id == pr.merchant_id)
        ).scalar_one_or_none()
        if not merchant:
            logger.error(
                "_charge_split_installment: split %s — merchant %s not found",
                split_id,
                pr.merchant_id,
            )
            return False

        currency = pr.currency.lower() if pr.currency else "usd"
        idempotency_key = f"{pr.payment_request_id}_split_{split.sequence or split.id}"

        # Submit charge
        try:
            provider_config, charge_result_obj = _run_provider_charge(
                merchant=merchant,
                db=db,
                pm=pm,
                amount_dollars=amount_dollars,
                currency=currency,
                idempotency_key=idempotency_key,
                customer_id=pr_customer.customer_id,
            )
        except Exception as charge_exc:
            logger.error(
                "_charge_split_installment: provider error for split %s: %s",
                split_id,
                charge_exc,
                exc_info=True,
            )
            return False

        charge_result = {
            "status": charge_result_obj.status,
            "transaction_id": charge_result_obj.transaction_id,
            "amount": charge_result_obj.amount,
            "currency": charge_result_obj.currency,
            "raw": charge_result_obj.raw_response,
        }

        if charge_result.get("status") == "succeeded":
            now = datetime.now(timezone.utc)
            split.paid_date = now

            txn_id = charge_result.get("transaction_id") or generate_secure_id(
                prepend="txn", length=20
            )
            transaction = create_transaction(
                db,
                txn_id=txn_id,
                txn_amount=float(txn_amount_cents),  # cents, consistent with services.py
                currency=currency,
                txn_status=TransactionStatusTypes.PAID,
                txn_type=pm.method or "card",
                txn_source="scheduled_split",
                category=TransactionCategories.CHARGE,
                transaction_type=TransactionTypes.PAYMENT_TERMINAL,
                txn_metadata=charge_result,
                payment_request_id=pr.id,
                payment_request_split_id=split.id,
                merchant_id=merchant.id,
                customer_id=pr_customer.customer_id,
                payment_method_id=pm.id,
            )
            _persist_provider_transaction(db, transaction, provider_config, charge_result)

            # Promote parent PR to PAID if every installment is now settled
            all_splits = db.execute(
                select(SplitPaymentRequests).where(
                    SplitPaymentRequests.payment_request_id == pr.id
                )
            ).scalars().all()
            all_splits_paid = all(s.paid_date is not None for s in all_splits)
            if all_splits_paid:
                pr.status = PaymentRequestStatusTypes.PAID
                logger.info(
                    "_charge_split_installment: PR %s — all splits paid → PAID",
                    pr.payment_request_id,
                )

            # Update invoice status: PARTIALLY_PAID while installments remain, PAID when all done
            try:
                from src.apps.invoices.models.invoice import Invoice
                from src.core.utils.enums import InvoiceStatusTypes
                from src.apps.invoices.services.invoice_services import _resolve_split_invoice_status

                inv_stmt = select(Invoice).where(
                    Invoice.payment_request_id == pr.id,
                    Invoice.deleted_at.is_(None),
                )
                invoice = db.execute(inv_stmt).scalar_one_or_none()
                if invoice is not None:
                    new_inv_status = _resolve_split_invoice_status(db, pr)
                    invoice.status = new_inv_status
                    if new_inv_status == InvoiceStatusTypes.PAID:
                        invoice.paid_date = invoice.paid_date or now
                    db.flush()
                    logger.info(
                        "_charge_split_installment: invoice for PR %s → %s",
                        pr.id,
                        new_inv_status.name,
                    )
            except Exception as _inv_exc:
                logger.warning(
                    "_charge_split_installment: could not update invoice status for PR %s: %s",
                    pr.id,
                    _inv_exc,
                )

            logger.info(
                "_charge_split_installment: split %s → PAID (txn_id=%s, amount=%.2f %s)",
                split_id,
                txn_id,
                amount_dollars,
                currency.upper(),
            )
            return True
        else:
            _fail_txn_id = charge_result.get("transaction_id") or generate_secure_id(prepend="txn", length=20)
            fail_transaction = create_transaction(
                db,
                txn_id=_fail_txn_id,
                txn_amount=float(txn_amount_cents),
                currency=currency,
                txn_status=TransactionStatusTypes.FAILED,
                txn_type=pm.method or "card",
                txn_source="scheduled_split",
                category=TransactionCategories.CHARGE,
                transaction_type=TransactionTypes.PAYMENT_TERMINAL,
                txn_metadata=charge_result,
                payment_request_id=pr.id,
                payment_request_split_id=split.id,
                merchant_id=merchant.id,
                customer_id=pr_customer.customer_id,
                payment_method_id=pm.id,
            )
            _persist_provider_transaction(db, fail_transaction, provider_config, charge_result)
            logger.warning(
                "_charge_split_installment: split %s charge FAILED (provider status=%s)",
                split_id,
                charge_result.get("status"),
            )
            return False
