from datetime import timedelta
from urllib.parse import urlencode

from django.db.models import Count, Max, Min
from django.db.models.functions import TruncDay
from django.utils import timezone

from rest_framework import viewsets
from rest_framework.decorators import action
from rest_framework.pagination import PageNumberPagination
from rest_framework.response import Response
from drf_spectacular.utils import extend_schema, OpenApiParameter

from apps.core.tenancy import TenantScopedMixin, get_active_organization

from .models import AccessLog, Hit, Site
from .serializers import AccessLogSerializer, HitSerializer, SiteSerializer
from .metadata import fetch_site_metadata


class AccessLogPagination(PageNumberPagination):
    page_size = 50
    page_size_query_param = "page_size"
    max_page_size = 500


def _top(qs, field, limit=10):
    return list(
        qs.exclude(**{f"{field}__in": ["", None]})
          .values(field)
          .annotate(count=Count("id"))
          .order_by("-count")[:limit]
    )


def _aggregate(qs):
    total = qs.count()
    unique_visitors = qs.exclude(visitor_hash="").values("visitor_hash").distinct().count()

    by_country = _top(qs, "country")
    by_country_code = _top(qs, "country_code")
    by_device = _top(qs, "device")
    by_os = _top(qs, "os")
    by_browser = _top(qs, "browser")
    by_referer = _top(qs, "referer_host")
    by_utm_source = _top(qs, "utm_source")
    by_utm_campaign = _top(qs, "utm_campaign")

    # Tags live in JSON list — aggregate manually.
    tag_counter: dict[str, int] = {}
    for row in qs.exclude(tags=[]).values_list("tags", flat=True):
        if not row:
            continue
        for t in row:
            if not t:
                continue
            tag_counter[t] = tag_counter.get(t, 0) + 1
    by_tag = sorted(
        ({"tag": k, "count": v} for k, v in tag_counter.items()),
        key=lambda r: r["count"],
        reverse=True,
    )[:20]

    by_resource_type = list(
        qs.values("resource_type")
          .annotate(count=Count("id"))
          .order_by("-count")
    )

    timeseries = list(
        qs.annotate(day=TruncDay("created_at"))
          .values("day")
          .annotate(count=Count("id"))
          .order_by("day")
    )

    return {
        "total_clicks": total,
        "unique_visitors": unique_visitors,
        "by_resource_type": by_resource_type,
        "by_country": by_country,
        "by_country_code": by_country_code,
        "by_device": by_device,
        "by_os": by_os,
        "by_browser": by_browser,
        "by_referer": by_referer,
        "by_utm_source": by_utm_source,
        "by_utm_campaign": by_utm_campaign,
        "by_tag": by_tag,
        "timeseries": timeseries,
    }


class AccessLogViewSet(viewsets.ReadOnlyModelViewSet):
    """Browse and aggregate analytics access logs."""
    queryset = AccessLog.objects.all()
    serializer_class = AccessLogSerializer
    pagination_class = AccessLogPagination

    def get_queryset(self):
        qs = super().get_queryset().filter(organization=get_active_organization(self.request))
        params = self.request.query_params

        rtype = params.get("resource_type")
        rid = params.get("resource_id")
        slug = params.get("resource_slug")
        country = params.get("country")
        device = params.get("device")
        utm = params.get("utm_source")
        days = params.get("days")
        date_from = params.get("date_from")
        date_to = params.get("date_to")

        if rtype:
            qs = qs.filter(resource_type=rtype)
        if rid:
            qs = qs.filter(resource_id=rid)
        if slug:
            qs = qs.filter(resource_slug=slug)
        if country:
            qs = qs.filter(country__iexact=country)
        if device:
            qs = qs.filter(device=device)
        if utm:
            qs = qs.filter(utm_source=utm)
        if days:
            try:
                cutoff = timezone.now() - timedelta(days=int(days))
                qs = qs.filter(created_at__gte=cutoff)
            except ValueError:
                pass
        if date_from:
            qs = qs.filter(created_at__date__gte=date_from)
        if date_to:
            qs = qs.filter(created_at__date__lte=date_to)
        return qs

    @extend_schema(
        parameters=[
            OpenApiParameter(name="resource_type", type=str, description="shorten_url | publication | event"),
            OpenApiParameter(name="resource_id", type=str),
            OpenApiParameter(name="resource_slug", type=str),
            OpenApiParameter(name="country", type=str),
            OpenApiParameter(name="device", type=str, description="mobile | tablet | desktop | bot"),
            OpenApiParameter(name="utm_source", type=str),
            OpenApiParameter(name="days", type=int, description="Restrict to last N days"),
            OpenApiParameter(name="date_from", type=str, description="YYYY-MM-DD"),
            OpenApiParameter(name="date_to", type=str, description="YYYY-MM-DD"),
        ]
    )
    def list(self, request, *args, **kwargs):
        return super().list(request, *args, **kwargs)

    @action(detail=False, methods=["get"])
    def stats(self, request):
        # Aggregation scans the whole filtered AccessLog set; cache per
        # filter-combination for 60s to absorb repeated dashboard polls.
        from django.core.cache import cache

        org = get_active_organization(request)
        key = f"analytics_accesslog_stats:{org.id}:" + urlencode(sorted(request.query_params.items()))
        cached = cache.get(key)
        if cached is not None:
            return Response(cached)
        data = _aggregate(self.filter_queryset(self.get_queryset()))
        cache.set(key, data, 60)
        return Response(data)


def _date_filter(qs, params):
    """Apply days / date_from / date_to query params to a Hit queryset."""
    days = params.get("days")
    date_from = params.get("date_from")
    date_to = params.get("date_to")
    if days:
        try:
            qs = qs.filter(created_at__gte=timezone.now() - timedelta(days=int(days)))
        except ValueError:
            pass
    if date_from:
        qs = qs.filter(created_at__date__gte=date_from)
    if date_to:
        qs = qs.filter(created_at__date__lte=date_to)
    return qs


def _aggregate_hits(qs):
    pageviews = qs.filter(event_type=Hit.EVENT_PAGEVIEW)
    events = qs.filter(event_type=Hit.EVENT_CUSTOM)

    timeseries = list(
        pageviews.annotate(day=TruncDay("created_at"))
        .values("day")
        .annotate(count=Count("id"))
        .order_by("day")
    )

    # Per-session rollup (pageviews carrying a cookieless session id).
    total_pageviews = pageviews.count()
    sess = list(
        pageviews.exclude(session_id="")
        .values("session_id")
        .annotate(pv=Count("id"), first=Min("created_at"), last=Max("created_at"))
    )
    sessions = len(sess)
    bounces = sum(1 for s in sess if s["pv"] == 1)
    durations = [(s["last"] - s["first"]).total_seconds() for s in sess]

    return {
        "total_pageviews": total_pageviews,
        "total_events": events.count(),
        "unique_visitors": qs.exclude(visitor_hash="").values("visitor_hash").distinct().count(),
        "sessions": sessions,
        "pages_per_session": round(total_pageviews / sessions, 1) if sessions else 0,
        "bounce_rate": round(bounces / sessions * 100) if sessions else 0,
        "avg_session_seconds": round(sum(durations) / sessions) if sessions else 0,
        "by_path": _top(pageviews, "path", 15),
        "by_referer": _top(pageviews, "referer_host"),
        "by_country": _top(qs, "country"),
        "by_country_code": _top(qs, "country_code"),
        "by_device": _top(qs, "device"),
        "by_os": _top(qs, "os"),
        "by_browser": _top(qs, "browser"),
        "by_event": _top(events, "event_name"),
        "by_utm_source": _top(qs, "utm_source"),
        "by_utm_campaign": _top(qs, "utm_campaign"),
        "timeseries": timeseries,
    }


class HitPagination(PageNumberPagination):
    page_size = 50
    page_size_query_param = "page_size"
    max_page_size = 200


class SiteViewSet(TenantScopedMixin, viewsets.ModelViewSet):
    """CRUD for tracked websites. Gated by the `analytics` RBAC domain."""
    rbac_domain = "analytics"
    queryset = Site.objects.all()
    serializer_class = SiteSerializer
    pagination_class = HitPagination

    def perform_create(self, serializer):
        user = self.request.user
        site = serializer.save(
            created_by=user if user.is_authenticated else None,
            organization=get_active_organization(self.request),
        )
        meta = fetch_site_metadata(site.domain)
        if meta:
            for k, v in meta.items():
                setattr(site, k, v)
            site.save(update_fields=list(meta.keys()))

    @action(detail=True, methods=["post"])
    def refresh_metadata(self, request, pk=None):
        site = self.get_object()
        meta = fetch_site_metadata(site.domain)
        if meta:
            for k, v in meta.items():
                setattr(site, k, v)
            site.save(update_fields=list(meta.keys()))
        return Response(self.get_serializer(site).data)

    @extend_schema(
        parameters=[
            OpenApiParameter(name="days", type=int, description="Restrict to last N days"),
            OpenApiParameter(name="date_from", type=str, description="YYYY-MM-DD"),
            OpenApiParameter(name="date_to", type=str, description="YYYY-MM-DD"),
        ]
    )
    @action(detail=True, methods=["get"])
    def stats(self, request, pk=None):
        # Per-site hit aggregation (sessions, bounce, timeseries) is heavy;
        # cache per site + date-filter for 60s.
        from django.core.cache import cache

        site = self.get_object()
        key = f"analytics_site_stats:{site.pk}:" + urlencode(sorted(request.query_params.items()))
        cached = cache.get(key)
        if cached is not None:
            return Response(cached)
        data = _aggregate_hits(_date_filter(site.hits.all(), request.query_params))
        cache.set(key, data, 60)
        return Response(data)

    @action(detail=True, methods=["get"])
    def hits(self, request, pk=None):
        site = self.get_object()
        qs = _date_filter(site.hits.all(), request.query_params)
        etype = request.query_params.get("event_type")
        if etype:
            qs = qs.filter(event_type=etype)
        page = self.paginate_queryset(qs)
        ser = HitSerializer(page if page is not None else qs, many=True)
        if page is not None:
            return self.get_paginated_response(ser.data)
        return Response(ser.data)
