"""
Binary Tree Celery Tasks
========================

A. rebuild_binary_tree_metrics_task
   — Recalculate subtree counts, left/right volumes, level snapshots.

B. expire_and_forfeit_power_stream_task
   — Expire overdue bonuses; log forfeiture to PowerStreamForfeitureLog.

C. process_level_completion_task
   — Detect newly completed binary levels; issue rewards.

D. validate_binary_integrity_task
   — Validate no cyclic genealogy; report mismatches.

E. assign_binary_position_task
   — Async single-node placement (called after purchase).
"""

import logging
from celery import shared_task
from celery.exceptions import Retry
from celery.utils.log import get_task_logger

logger = get_task_logger(__name__)


# ── A. Binary Tree Metrics Rebuilder ─────────────────────────────────────────

@shared_task(bind=True, max_retries=2, default_retry_delay=60)
def rebuild_binary_tree_metrics_task(self, bundle_id=None, only_stale=True):
    """
    Full rebuild of BinaryTreeMetrics and BinaryLevelSnapshot for all
    (or bundle-specific) DistributorIDs.

    Runs nightly via Celery Beat.
    Also triggered by admin "Rebuild Metrics" action.
    """
    logger.info("[BINARY] rebuild_binary_tree_metrics_task START bundle=%s stale=%s", bundle_id, only_stale)
    try:
        from apps.business.binary_tree.services.metrics_service import compute_and_persist_all
        updated = compute_and_persist_all(bundle_id=bundle_id, only_stale=only_stale)
        logger.info("[BINARY] rebuild_binary_tree_metrics_task DONE updated=%d", updated)
        return {'updated': updated}
    except Exception as exc:
        logger.error("[BINARY] rebuild_binary_tree_metrics_task ERROR: %s", exc)
        try:
            raise self.retry(exc=exc)
        except Retry:
            raise


# ── B. Power Stream Expiry + Forfeiture ──────────────────────────────────────

@shared_task(bind=True, max_retries=2, default_retry_delay=60)
def expire_and_forfeit_power_stream_task(self):
    """
    1. Marks PENDING power stream bonuses whose expires_at < now as FORFEITED.
    2. Logs each forfeiture to PowerStreamForfeitureLog.
    3. Grace rule: checks whether user lost minimum 2 referrals this month;
       if so, marks bonuses FORFEITED (grace period ended).

    Runs daily via Celery Beat.
    """
    logger.info("[PS] expire_and_forfeit_power_stream_task START")
    try:
        from django.utils import timezone
        from apps.business.power_stream.models import PowerStreamBonusRecord
        from apps.business.binary_tree.models import PowerStreamForfeitureLog
        from apps.business.distributor.models import DistributorID

        now = timezone.now()
        overdue = list(
            PowerStreamBonusRecord.objects.filter(
                status=PowerStreamBonusRecord.Status.PENDING,
                expires_at__lt=now,
            ).select_related('earner')
        )

        forfeited_count = 0
        for record in overdue:
            try:
                # Determine if the user still has minimum 2 direct referrals
                user = record.earner
                max_direct = max(
                    (
                        DistributorID.objects.filter(
                            sponsor_distributor=d, is_active=True
                        ).count()
                        for d in DistributorID.objects.filter(user=user)
                    ),
                    default=0,
                )
                grace_expired = max_direct < 2  # lost minimum referrals

                record.status = PowerStreamBonusRecord.Status.FORFEITED
                record.save(update_fields=['status', 'updated_at'])

                PowerStreamForfeitureLog.objects.get_or_create(
                    bonus_record=record,
                    defaults={
                        'user': user,
                        'gross_amount': record.gross_amount,
                        'net_amount': record.net_amount,
                        'reason': (
                            'Grace period expired — minimum referrals not maintained.'
                            if grace_expired
                            else 'Bonus expired after 30-day claim window.'
                        ),
                        'grace_expired': grace_expired,
                    },
                )
                forfeited_count += 1
            except Exception as inner_exc:
                logger.error("[PS] Error forfeiting record %s: %s", record.id, inner_exc)

        logger.info("[PS] expire_and_forfeit_power_stream_task DONE forfeited=%d", forfeited_count)
        return {'forfeited': forfeited_count}
    except Exception as exc:
        logger.error("[PS] expire_and_forfeit_power_stream_task ERROR: %s", exc)
        try:
            raise self.retry(exc=exc)
        except Retry:
            raise


# ── C. Level Completion Worker ────────────────────────────────────────────────

@shared_task(bind=True, max_retries=3, default_retry_delay=30)
def process_level_completion_task(self, distributor_id):
    """
    For one DistributorID:
    1. Recalculate binary metrics + level snapshot.
    2. If new levels are unlocked (current_level > last_completed_level), issue rewards.

    Called after every new binary node placement (from assign_binary_position_task).
    """
    logger.info("[LEVEL] process_level_completion_task distributor_id=%s", distributor_id)
    try:
        from apps.business.distributor.models import DistributorID
        from apps.business.rewards.services.level_service import process_level_rewards
        from apps.business.binary_tree.services.metrics_service import (
            compute_metrics_for_distributor,
            upsert_metrics,
            compute_binary_level_snapshot,
        )

        try:
            distributor = DistributorID.objects.select_related('product__opportunity_bundle').get(pk=distributor_id)
        except DistributorID.DoesNotExist:
            logger.warning("[LEVEL] Distributor not found: %s", distributor_id)
            return

        metrics_data = compute_metrics_for_distributor(distributor)
        upsert_metrics(distributor, metrics_data)
        snapshot = compute_binary_level_snapshot(distributor)

        if snapshot and snapshot.current_level > snapshot.last_completed_level:
            result = process_level_rewards(distributor)
            logger.info(
                "[LEVEL] Rewards processed for %s: levels=%s total=%s",
                distributor_id, result.get('levels_rewarded'), result.get('total_reward'),
            )
            return result

        return {'levels_rewarded': [], 'total_reward': 0}
    except Exception as exc:
        logger.error("[LEVEL] process_level_completion_task ERROR %s: %s", distributor_id, exc)
        try:
            raise self.retry(exc=exc)
        except Retry:
            raise


# ── D. Binary Integrity Validator ────────────────────────────────────────────

@shared_task(bind=True, max_retries=1, default_retry_delay=120)
def validate_binary_integrity_task(self, bundle_id=None):
    """
    Scan binary tree for:
    • Cyclic parent chains
    • Parent/child position mismatches
    • Invalid binary_position values

    Runs weekly. Logs and returns a report.
    """
    logger.info("[INTEGRITY] validate_binary_integrity_task START bundle=%s", bundle_id)
    try:
        from apps.business.binary_tree.services.binary_service import verify_binary_integrity
        report = verify_binary_integrity(bundle_id=bundle_id)
        if report['issues_found']:
            logger.warning(
                "[INTEGRITY] Issues detected: cycles=%d mismatches=%d invalid=%d",
                len(report['cycles']),
                len(report['orphan_mismatches']),
                len(report['invalid_positions']),
            )
        else:
            logger.info("[INTEGRITY] No issues found. total_checked=%d", report['total_checked'])
        return report
    except Exception as exc:
        logger.error("[INTEGRITY] validate_binary_integrity_task ERROR: %s", exc)
        try:
            raise self.retry(exc=exc)
        except Retry:
            raise


# ── E. Async Binary Placement (called after purchase) ────────────────────────

@shared_task(bind=True, max_retries=3, default_retry_delay=15)
def assign_binary_position_task(self, distributor_id, sponsor_distributor_id):
    """
    Legacy sponsor-based binary placement task.

    Strict hybrid mode forbids sponsor-driven binary placement. This task is kept
    only to fail closed if an old caller still exists.
    """
    message = (
        "assign_binary_position_task is disabled. "
        "Use create_distributor_id() and company-root BFS placement only."
    )
    logger.warning(
        "[BINARY] %s distributor_id=%s sponsor_distributor_id=%s",
        message,
        distributor_id,
        sponsor_distributor_id,
    )
    raise RuntimeError(message)


def _enqueue_ancestor_level_checks(distributor, max_depth=20):
    """
    Walk up the binary_parent chain and enqueue process_level_completion_task
    for each ancestor so their level snapshots are refreshed.
    """
    from apps.business.distributor.models import DistributorID
    visited = set()
    current_id = distributor.binary_parent_distributor_id
    depth = 0
    while current_id and depth < max_depth:
        if current_id in visited:
            break
        visited.add(current_id)
        process_level_completion_task.delay(str(current_id))
        current_id = (
            DistributorID.objects.filter(pk=current_id)
            .values_list('binary_parent_distributor_id', flat=True)
            .first()
        )
        depth += 1
