from datetime import datetime, timedelta
from decimal import Decimal

from django.db.models import Count, Min, Q
from django.utils import timezone

from apps.business.distributor.models import DistributorID
from apps.telemetry.models import (
    TelemetryAlert,
    TelemetryCohortSnapshot,
    TelemetryEvent,
    TelemetryFunnelSnapshot,
    TelemetryKPIAggregate,
    TelemetryRetentionSnapshot,
)


EVENT_DOMAIN_MAP = {
    'recommendation_impression': TelemetryEvent.Domain.COMMERCE,
    'recommendation_click': TelemetryEvent.Domain.COMMERCE,
    'reorder_click': TelemetryEvent.Domain.COMMERCE,
    'bundle_upgrade_click': TelemetryEvent.Domain.COMMERCE,
    'marketplace_product_click': TelemetryEvent.Domain.COMMERCE,
    'notification_open': TelemetryEvent.Domain.COMMERCE,
    'widget_interaction': TelemetryEvent.Domain.COMMERCE,
    'wallet_usage': TelemetryEvent.Domain.COMMERCE,
    'checkout_start': TelemetryEvent.Domain.COMMERCE,
    'checkout_complete': TelemetryEvent.Domain.COMMERCE,
    'page_view': TelemetryEvent.Domain.COMMERCE,
    'api_timing': TelemetryEvent.Domain.OPERATIONAL,
    'frontend_perf': TelemetryEvent.Domain.OPERATIONAL,
    'client_error': TelemetryEvent.Domain.OPERATIONAL,
    'notification_delivery_failure': TelemetryEvent.Domain.OPERATIONAL,
}


SENSITIVE_KEYS = {
    'email', 'mobile', 'phone', 'pan', 'aadhaar', 'token', 'access_token', 'refresh_token', 'password',
}


def sanitize_metadata(payload: dict) -> dict:
    if not isinstance(payload, dict):
        return {}
    cleaned = {}
    for key, value in payload.items():
        k = str(key).strip().lower()
        if k in SENSITIVE_KEYS:
            continue
        cleaned[str(key)] = value
    return cleaned


def should_sample(event_type: str, metadata: dict) -> bool:
    if event_type in {'api_timing', 'frontend_perf'}:
        duration = float(metadata.get('durationMs') or metadata.get('duration_ms') or 0)
        # Keep all slow events, sample only the noisy fast-path.
        if duration > 1200:
            return False
        return True
    return False


def ingest_events(user, events, source='frontend'):
    now = timezone.now()
    retention_cutoff = now + timedelta(days=90)
    ingested = 0
    sampled = 0
    dropped = 0
    to_create = []

    for item in events:
        event_type = item['event_type']
        event_domain = EVENT_DOMAIN_MAP.get(event_type, TelemetryEvent.Domain.COMMERCE)
        metadata = sanitize_metadata(item.get('metadata') or {})

        if should_sample(event_type, metadata):
            sampled += 1
            # keep only 1 in 3 sampled events
            if (now.microsecond + sampled) % 3 != 0:
                continue

        occurred_at = item.get('timestamp') or now
        try:
            to_create.append(
                TelemetryEvent(
                    event_type=event_type,
                    event_domain=event_domain,
                    occurred_at=occurred_at,
                    user=user if getattr(user, 'is_authenticated', False) else None,
                    session_id=(item.get('session') or '')[:128],
                    device=(item.get('device') or '')[:64],
                    page=(item.get('page') or '')[:255],
                    source=(source or 'frontend')[:32],
                    metadata=metadata,
                    client_event_id=(item.get('client_event_id') or '')[:128],
                    retention_expires_at=retention_cutoff,
                )
            )
        except Exception:
            dropped += 1

    if to_create:
        TelemetryEvent.objects.bulk_create(to_create, batch_size=1000)
        ingested = len(to_create)

    return {
        'ingested': ingested,
        'sampled': sampled,
        'dropped': dropped,
    }


def _period_bounds(period_type: str, ref_date):
    if period_type == TelemetryKPIAggregate.PeriodType.DAY:
        start = ref_date
        end = ref_date
    elif period_type == TelemetryKPIAggregate.PeriodType.WEEK:
        start = ref_date - timedelta(days=ref_date.weekday())
        end = start + timedelta(days=6)
    else:
        start = ref_date.replace(day=1)
        if start.month == 12:
            month_end = start.replace(year=start.year + 1, month=1, day=1)
        else:
            month_end = start.replace(month=start.month + 1, day=1)
        end = month_end - timedelta(days=1)
    return start, end


def aggregate_kpis(period_type: str, ref_date=None):
    ref_date = ref_date or timezone.localdate()
    period_start, period_end = _period_bounds(period_type, ref_date)
    start_dt = timezone.make_aware(datetime.combine(period_start, datetime.min.time()))
    end_dt = timezone.make_aware(datetime.combine(period_end, datetime.max.time()))

    qs = TelemetryEvent.objects.filter(occurred_at__gte=start_dt, occurred_at__lte=end_dt)

    for domain in TelemetryEvent.Domain.values:
        domain_qs = qs.filter(event_domain=domain)
        total_events = domain_qs.count()
        unique_users = domain_qs.exclude(user__isnull=True).values('user_id').distinct().count()

        rec_impressions = domain_qs.filter(event_type='recommendation_impression').count()
        rec_clicks = domain_qs.filter(event_type='recommendation_click').count()
        rec_ctr = round((rec_clicks / rec_impressions) * 100, 2) if rec_impressions else 0

        checkout_start = domain_qs.filter(event_type='checkout_start').count()
        checkout_complete = domain_qs.filter(event_type='checkout_complete').count()
        checkout_conversion = round((checkout_complete / checkout_start) * 100, 2) if checkout_start else 0

        wallet_usage = domain_qs.filter(event_type='wallet_usage').count()
        wallet_active_users = domain_qs.filter(event_type='wallet_usage').exclude(user__isnull=True).values('user_id').distinct().count()

        api_failures = domain_qs.filter(event_type='api_timing').filter(
            Q(metadata__status='500') | Q(metadata__status='502') | Q(metadata__status='503')
        ).count()
        slow_api = domain_qs.filter(event_type='api_timing').filter(
            Q(metadata__durationMs__gt=1200) | Q(metadata__duration_ms__gt=1200)
        ).count()
        client_errors = domain_qs.filter(event_type='client_error').count()

        metrics = {
            'total_events': total_events,
            'unique_users': unique_users,
            'recommendation_impressions': rec_impressions,
            'recommendation_clicks': rec_clicks,
            'recommendation_ctr': rec_ctr,
            'checkout_start': checkout_start,
            'checkout_complete': checkout_complete,
            'checkout_conversion': checkout_conversion,
            'wallet_usage': wallet_usage,
            'wallet_active_users': wallet_active_users,
            'api_failures': api_failures,
            'slow_api': slow_api,
            'client_errors': client_errors,
        }

        TelemetryKPIAggregate.objects.update_or_create(
            period_type=period_type,
            period_start=period_start,
            event_domain=domain,
            defaults={
                'period_end': period_end,
                'metrics': metrics,
            },
        )


def compute_cohort_snapshots(snapshot_date=None):
    snapshot_date = snapshot_date or timezone.localdate()
    start_dt = timezone.now() - timedelta(days=30)
    events = TelemetryEvent.objects.filter(occurred_at__gte=start_dt).exclude(user__isnull=True)

    user_stats = {}
    for row in events.values('user_id', 'event_type').annotate(c=Count('id')):
        user_id = row['user_id']
        user_stats.setdefault(user_id, {})[row['event_type']] = row['c']

    dist_counts = dict(
        DistributorID.objects.values('user_id').annotate(c=Count('id')).values_list('user_id', 'c')
    )

    labels = {
        'commerce_heavy': 0,
        'network_heavy': 0,
        'dormant_users': 0,
        'repeat_purchasers': 0,
        'wallet_active_users': 0,
        'bundle_upgrade_users': 0,
    }

    active_recent = set(
        TelemetryEvent.objects.filter(occurred_at__gte=timezone.now() - timedelta(days=14))
        .exclude(user__isnull=True)
        .values_list('user_id', flat=True)
        .distinct()
    )

    for user_id, stats in user_stats.items():
        commerce_score = sum(stats.get(k, 0) for k in (
            'marketplace_product_click',
            'checkout_start',
            'checkout_complete',
            'wallet_usage',
            'recommendation_click',
        ))
        network_score = dist_counts.get(user_id, 0)

        if commerce_score >= max(5, network_score * 2):
            labels['commerce_heavy'] += 1
        if network_score > max(1, commerce_score):
            labels['network_heavy'] += 1
        if user_id not in active_recent:
            labels['dormant_users'] += 1
        if stats.get('checkout_complete', 0) >= 2:
            labels['repeat_purchasers'] += 1
        if stats.get('wallet_usage', 0) >= 1:
            labels['wallet_active_users'] += 1
        if stats.get('bundle_upgrade_click', 0) >= 1:
            labels['bundle_upgrade_users'] += 1

    for cohort_name, count in labels.items():
        TelemetryCohortSnapshot.objects.update_or_create(
            snapshot_date=snapshot_date,
            cohort_name=cohort_name,
            defaults={
                'users_count': count,
                'metrics': {'window_days': 30},
            },
        )


def compute_retention_snapshots(max_days=30):
    today = timezone.localdate()
    first_seen = (
        TelemetryEvent.objects.exclude(user__isnull=True)
        .values('user_id')
        .annotate(first_seen=Min('occurred_at'))
    )

    cohort_map = {}
    for item in first_seen:
        dt = item['first_seen']
        if not dt:
            continue
        cohort_map.setdefault(dt.date(), set()).add(item['user_id'])

    for cohort_date, users in cohort_map.items():
        cohort_size = len(users)
        if cohort_size == 0:
            continue
        for day in (1, 7, 30):
            target_date = cohort_date + timedelta(days=day)
            if target_date > today:
                continue
            active_users = TelemetryEvent.objects.filter(
                user_id__in=users,
                occurred_at__date=target_date,
            ).values('user_id').distinct().count()
            rate = Decimal('0.00')
            if cohort_size:
                rate = Decimal(str(round((active_users / cohort_size) * 100, 2)))

            TelemetryRetentionSnapshot.objects.update_or_create(
                cohort_date=cohort_date,
                retention_day=day,
                event_domain=TelemetryEvent.Domain.COMMERCE,
                defaults={
                    'cohort_size': cohort_size,
                    'active_users': active_users,
                    'retention_rate': rate,
                },
            )


def compute_funnel_snapshots(snapshot_date=None):
    snapshot_date = snapshot_date or timezone.localdate()
    start_dt = timezone.make_aware(datetime.combine(snapshot_date, datetime.min.time()))
    end_dt = timezone.make_aware(datetime.combine(snapshot_date, datetime.max.time()))
    day_qs = TelemetryEvent.objects.filter(occurred_at__gte=start_dt, occurred_at__lte=end_dt)

    commerce_funnel = [
        ('Dashboard Visit', 'page_view'),
        ('Product Click', 'marketplace_product_click'),
        ('Cart', 'widget_interaction'),
        ('Checkout', 'checkout_start'),
        ('Purchase', 'checkout_complete'),
        ('Repeat Purchase', 'reorder_click'),
    ]
    wallet_funnel = [
        ('Wallet Credit', 'wallet_usage'),
        ('Marketplace Usage', 'marketplace_product_click'),
        ('Reorder', 'reorder_click'),
    ]

    for funnel_name, definition in (
        ('commerce_conversion', commerce_funnel),
        ('wallet_reorder', wallet_funnel),
    ):
        baseline = 0
        for idx, (stage_name, event_type) in enumerate(definition):
            users_count = day_qs.filter(event_type=event_type).exclude(user__isnull=True).values('user_id').distinct().count()
            if idx == 0:
                baseline = max(users_count, 1)
            conversion = round((users_count / baseline) * 100, 2) if baseline else 0
            TelemetryFunnelSnapshot.objects.update_or_create(
                snapshot_date=snapshot_date,
                funnel_name=funnel_name,
                stage_order=idx + 1,
                event_domain=TelemetryEvent.Domain.COMMERCE,
                defaults={
                    'stage_name': stage_name,
                    'users_count': users_count,
                    'conversion_rate': Decimal(str(conversion)),
                },
            )


def _create_alert(alert_type, severity, metric_value, threshold_value, context):
    existing = TelemetryAlert.objects.filter(
        alert_type=alert_type,
        status=TelemetryAlert.Status.OPEN,
        detected_at__gte=timezone.now() - timedelta(hours=6),
    ).first()
    if existing:
        return existing

    return TelemetryAlert.objects.create(
        alert_type=alert_type,
        severity=severity,
        detected_at=timezone.now(),
        metric_value=Decimal(str(metric_value)),
        threshold_value=Decimal(str(threshold_value)),
        context=context,
    )


def detect_anomalies():
    now = timezone.now()
    last_hour = now - timedelta(hours=1)
    baseline_since = now - timedelta(days=1)

    events_last_hour = TelemetryEvent.objects.filter(occurred_at__gte=last_hour)
    events_baseline = TelemetryEvent.objects.filter(occurred_at__gte=baseline_since, occurred_at__lt=last_hour)

    failures_last_hour = events_last_hour.filter(event_type='api_timing').filter(
        Q(metadata__status='500') | Q(metadata__status='502') | Q(metadata__status='503')
    ).count()
    baseline_failures = events_baseline.filter(event_type='api_timing').filter(
        Q(metadata__status='500') | Q(metadata__status='502') | Q(metadata__status='503')
    ).count()
    baseline_failure_hourly = baseline_failures / 23 if baseline_failures else 0
    if failures_last_hour > max(10, baseline_failure_hourly * 2):
        _create_alert('api_failure_spike', TelemetryAlert.Severity.CRITICAL, failures_last_hour, baseline_failure_hourly, {'window': '1h'})

    starts = events_last_hour.filter(event_type='checkout_start').count()
    completes = events_last_hour.filter(event_type='checkout_complete').count()
    current_conv = (completes / starts) if starts else 0

    b_starts = events_baseline.filter(event_type='checkout_start').count()
    b_completes = events_baseline.filter(event_type='checkout_complete').count()
    baseline_conv = (b_completes / b_starts) if b_starts else 0
    if starts >= 20 and baseline_conv > 0 and current_conv < (baseline_conv * 0.7):
        _create_alert('checkout_drop_spike', TelemetryAlert.Severity.WARNING, current_conv * 100, baseline_conv * 100, {'window': '1h'})

    wallet_last_hour = events_last_hour.filter(event_type='wallet_usage').count()
    wallet_baseline = events_baseline.filter(event_type='wallet_usage').count() / 23 if events_baseline.exists() else 0
    if wallet_last_hour > max(30, wallet_baseline * 3):
        _create_alert('wallet_anomaly_spike', TelemetryAlert.Severity.WARNING, wallet_last_hour, wallet_baseline, {'window': '1h'})

    imp = events_last_hour.filter(event_type='recommendation_impression').count()
    clk = events_last_hour.filter(event_type='recommendation_click').count()
    ctr = (clk / imp) if imp else 0
    b_imp = events_baseline.filter(event_type='recommendation_impression').count()
    b_clk = events_baseline.filter(event_type='recommendation_click').count()
    b_ctr = (b_clk / b_imp) if b_imp else 0
    if imp >= 50 and b_ctr > 0 and ctr < (b_ctr * 0.5):
        _create_alert('recommendation_failure_spike', TelemetryAlert.Severity.WARNING, ctr * 100, b_ctr * 100, {'window': '1h'})

    notif_failures = events_last_hour.filter(event_type='notification_delivery_failure').count()
    if notif_failures > 10:
        _create_alert('notification_delivery_failures', TelemetryAlert.Severity.WARNING, notif_failures, 10, {'window': '1h'})


def compact_and_apply_retention():
    now = timezone.now()
    TelemetryEvent.objects.filter(retention_expires_at__lt=now).delete()
    # Keep detailed alerts for 180 days; historical KPI tables remain for long-term trends.
    TelemetryAlert.objects.filter(detected_at__lt=now - timedelta(days=180), status=TelemetryAlert.Status.RESOLVED).delete()
