from __future__ import annotations

from collections import defaultdict
from datetime import timedelta
from urllib.parse import urlencode

from django.utils import timezone
from rest_framework import permissions, status
from rest_framework.response import Response
from rest_framework.views import APIView

from apps.accounts.models import User


class UserReferralOverviewView(APIView):
    permission_classes = [permissions.IsAuthenticated]

    def get(self, request):
        user = request.user
        max_depth = self._to_int(request.query_params.get('max_depth'), default=4, min_value=1, max_value=8)
        max_nodes = self._to_int(request.query_params.get('max_nodes'), default=250, min_value=20, max_value=1000)

        descendants, parent_children = self._collect_referral_graph(root_user=user, max_depth=max_depth, max_nodes=max_nodes)
        direct_referral_ids = parent_children.get(str(user.id), [])

        direct_referrals_qs = (
            User.objects.filter(id__in=direct_referral_ids)
            .order_by('-date_joined')
            .values('id', 'first_name', 'last_name', 'mobile', 'email', 'referral_code', 'date_joined')
        )
        direct_referrals = [
            {
                'id': str(row['id']),
                'name': self._display_name(row),
                'mobile': row.get('mobile') or '',
                'email': row.get('email') or '',
                'referral_code': row.get('referral_code') or '',
                'joined_at': row['date_joined'],
                'level': 1,
                'direct_referrals': len(parent_children.get(str(row['id']), [])),
            }
            for row in direct_referrals_qs
        ]

        levels = defaultdict(int)
        for row in descendants.values():
            levels[row['level']] += 1

        level_summary = [
            {'level': level, 'count': levels[level]}
            for level in sorted(levels.keys())
        ]

        referral_link = self._build_referral_link(request, user.referral_code)
        joined_recently_cutoff = timezone.now() - timedelta(days=30)

        response_payload = {
            'summary': {
                'direct_referrals': len(direct_referral_ids),
                'total_network': len(descendants),
                'max_depth_loaded': max(levels.keys()) if levels else 0,
                'new_referrals_30d': sum(1 for row in descendants.values() if row['date_joined'] >= joined_recently_cutoff),
            },
            'share': {
                'referral_code': user.referral_code or '',
                'referral_link': referral_link,
                'whatsapp_share_link': self._build_whatsapp_link(referral_link),
            },
            'direct_referrals': direct_referrals,
            'level_summary': level_summary,
            'tree': self._build_tree(user, descendants, parent_children, max_depth=max_depth),
            'meta': {
                'max_depth_requested': max_depth,
                'max_nodes_requested': max_nodes,
                'loaded_nodes': len(descendants),
                'truncated': len(descendants) >= max_nodes,
            },
        }

        return Response({'success': True, 'data': response_payload}, status=status.HTTP_200_OK)

    @staticmethod
    def _to_int(value, default: int, min_value: int, max_value: int) -> int:
        try:
            parsed = int(value)
        except (TypeError, ValueError):
            return default
        return max(min_value, min(max_value, parsed))

    @staticmethod
    def _display_name(row: dict) -> str:
        first_name = (row.get('first_name') or '').strip()
        last_name = (row.get('last_name') or '').strip()
        full_name = f'{first_name} {last_name}'.strip()
        return full_name or (row.get('mobile') or row.get('email') or 'User')

    def _collect_referral_graph(self, root_user: User, max_depth: int, max_nodes: int):
        descendants: dict[str, dict] = {}
        parent_children: dict[str, list[str]] = defaultdict(list)

        parent_ids = [root_user.id]
        current_depth = 1

        while parent_ids and current_depth <= max_depth and len(descendants) < max_nodes:
            rows = list(
                User.objects.filter(referred_by_id__in=parent_ids)
                .order_by('date_joined', 'id')
                .values('id', 'first_name', 'last_name', 'mobile', 'email', 'referral_code', 'date_joined', 'referred_by_id')
            )

            next_parent_ids = []
            for row in rows:
                user_id = str(row['id'])
                if user_id in descendants:
                    continue
                if len(descendants) >= max_nodes:
                    break

                parent_id = str(row['referred_by_id']) if row.get('referred_by_id') else ''
                descendants[user_id] = {
                    'id': user_id,
                    'parent_id': parent_id,
                    'name': self._display_name(row),
                    'mobile': row.get('mobile') or '',
                    'email': row.get('email') or '',
                    'referral_code': row.get('referral_code') or '',
                    'date_joined': row['date_joined'],
                    'level': current_depth,
                }
                parent_children[parent_id].append(user_id)
                next_parent_ids.append(row['id'])

            parent_ids = next_parent_ids
            current_depth += 1

        return descendants, parent_children

    def _build_tree(self, root_user: User, descendants: dict[str, dict], parent_children: dict[str, list[str]], max_depth: int):
        root_id = str(root_user.id)

        def node_for(user_id: str, level: int):
            if user_id == root_id:
                root_name = self._display_name(
                    {
                        'first_name': getattr(root_user, 'first_name', ''),
                        'last_name': getattr(root_user, 'last_name', ''),
                        'mobile': getattr(root_user, 'mobile', ''),
                        'email': getattr(root_user, 'email', ''),
                    }
                )
                base = {
                    'id': root_id,
                    'name': root_name,
                    'mobile': root_user.mobile or '',
                    'email': root_user.email or '',
                    'referral_code': root_user.referral_code or '',
                    'joined_at': root_user.date_joined,
                    'level': 0,
                }
            else:
                row = descendants[user_id]
                base = {
                    'id': row['id'],
                    'name': row['name'],
                    'mobile': row['mobile'],
                    'email': row['email'],
                    'referral_code': row['referral_code'],
                    'joined_at': row['date_joined'],
                    'level': row['level'],
                }

            child_ids = parent_children.get(user_id, [])
            base['direct_referrals'] = len(child_ids)
            if level >= max_depth:
                base['children'] = []
                return base

            base['children'] = [node_for(child_id, level + 1) for child_id in child_ids]
            return base

        return node_for(root_id, 0)

    @staticmethod
    def _build_referral_link(request, referral_code: str | None) -> str:
        if not referral_code:
            return ''
        frontend_base = request.headers.get('Origin') or 'https://www.truewaveindia.com'
        base = frontend_base.rstrip('/')
        query = urlencode({'ref': referral_code})
        return f'{base}/register?{query}'

    @staticmethod
    def _build_whatsapp_link(referral_link: str) -> str:
        if not referral_link:
            return ''
        message = f'Join my TrueWave network: {referral_link}'
        return f'https://wa.me/?{urlencode({"text": message})}'
