from rest_framework import viewsets, status, permissions
from rest_framework.decorators import action
from rest_framework.response import Response
from rest_framework.views import APIView
from rest_framework.parsers import MultiPartParser, FormParser
from rest_framework.throttling import ScopedRateThrottle
from drf_spectacular.utils import extend_schema, OpenApiParameter
from django.db.models import Count, OuterRef, Subquery, IntegerField, Prefetch, Max
from django.contrib.auth import get_user_model

from apps.core.pagination import StandardResultsSetPagination
from apps.core.tenancy import get_active_organization
from .link_preview import fetch_link_preview, LinkPreviewError

from .models import (
    Conversation,
    ConversationParticipant,
    Message,
    MessageReaction,
    MessageKeyWrap,
    MessageAttachment,
    UserKey,
)
from .serializers import (
    ConversationSerializer,
    ConversationCreateSerializer,
    MessageSerializer,
    MessageCreateSerializer,
    MessageReactionSerializer,
    MessageAttachmentSerializer,
    MessageAttachmentUploadSerializer,
    UserKeySerializer,
)

User = get_user_model()


class UserKeyViewSet(viewsets.ViewSet):
    """Manage X25519 public keys for E2E chat."""
    permission_classes = [permissions.IsAuthenticated]

    @action(detail=False, methods=['get', 'put'], url_path='me')
    def me(self, request):
        if request.method == 'GET':
            key = UserKey.objects.filter(user=request.user).first()
            if not key:
                return Response({'public_key': None})
            return Response(UserKeySerializer(key).data)
        public_key = request.data.get('public_key')
        if not public_key or len(public_key) > 128:
            return Response({'error': 'public_key is required'}, status=status.HTTP_400_BAD_REQUEST)
        key, _ = UserKey.objects.update_or_create(
            user=request.user,
            defaults={'public_key': public_key},
        )
        return Response(UserKeySerializer(key).data)

    @action(detail=False, methods=['post'], url_path='me/claim')
    def claim_bootstrap(self, request):
        """Return the server-escrowed private key for the requesting user, then
        wipe it from the server. Single-use: subsequent calls return 404.
        Client must persist the returned `private_key` in IndexedDB before
        re-calling, or rotate to a fresh client-generated key via PUT /me/.
        """
        key = UserKey.objects.filter(user=request.user).first()
        if not key or not key.private_key_escrow or key.escrow_claimed:
            return Response(
                {'error': 'No escrowed private key for this user.'},
                status=status.HTTP_404_NOT_FOUND,
            )
        priv = key.private_key_escrow
        key.private_key_escrow = ''
        key.escrow_claimed = True
        key.save(update_fields=['private_key_escrow', 'escrow_claimed', 'rotated_at'])
        return Response({
            'public_key': key.public_key,
            'private_key': priv,
        })

    @extend_schema(
        description='Bulk fetch public keys for a set of users.',
        parameters=[OpenApiParameter(name='user_ids', type=str, description='Comma-separated UUIDs')],
    )
    def list(self, request):
        ids_param = request.query_params.get('user_ids', '')
        if not ids_param:
            return Response([])
        ids = [i for i in ids_param.split(',') if i]
        organization = get_active_organization(request)
        keys = UserKey.objects.filter(
            user_id__in=ids, user__organizations=organization,
        ).select_related('user')
        return Response(UserKeySerializer(keys, many=True).data)


class ConversationViewSet(viewsets.ModelViewSet):
    serializer_class = ConversationSerializer
    permission_classes = [permissions.IsAuthenticated]
    pagination_class = StandardResultsSetPagination

    def get_queryset(self):
        user = self.request.user
        organization = get_active_organization(self.request)
        unread_sq = (
            Message.objects
            .filter(conversation=OuterRef('pk'), is_read=False)
            .exclude(sender=user)
            .order_by()
            .values('conversation')
            .annotate(c=Count('id'))
            .values('c')[:1]
        )
        return (
            Conversation.objects
            .filter(participants=user, organization=organization)
            .select_related('created_by')
            .prefetch_related(
                Prefetch(
                    'conversationparticipant_set',
                    queryset=ConversationParticipant.objects.select_related('user', 'user__chat_key'),
                ),
            )
            .annotate(
                unread_count_anno=Subquery(unread_sq, output_field=IntegerField()),
                last_message_at_anno=Max('messages__created_at'),
            )
            .distinct()
            .order_by('-last_message_at_anno', '-pk')
        )

    def get_serializer_class(self):
        if self.action == 'create':
            return ConversationCreateSerializer
        return ConversationSerializer

    def perform_create(self, serializer):
        serializer.save(created_by=self.request.user, organization=get_active_organization(self.request))

    @extend_schema(
        description='Get or create a direct 1:1 conversation with another user.',
    )
    @action(detail=False, methods=['post'])
    def direct(self, request):
        user_id = request.data.get('user_id')
        if not user_id:
            return Response({'error': 'user_id is required'}, status=status.HTTP_400_BAD_REQUEST)

        organization = get_active_organization(request)

        existing = (
            Conversation.objects
            .filter(is_group=False, participants=request.user, organization=organization)
            .filter(participants__id=user_id)
            .distinct()
        )
        if existing.exists():
            conversation = existing.first()
        else:
            try:
                other_user = User.objects.get(id=user_id, organizations=organization)
            except User.DoesNotExist:
                return Response({'error': 'User not found'}, status=status.HTTP_404_NOT_FOUND)
            conversation = Conversation.objects.create(
                created_by=request.user, is_group=False, organization=organization,
            )
            ConversationParticipant.objects.create(
                conversation=conversation, user=request.user, is_admin=True,
            )
            ConversationParticipant.objects.create(conversation=conversation, user=other_user)

        serializer = ConversationSerializer(conversation, context={'request': request})
        return Response(serializer.data)

    @extend_schema(description='Paginated messages in a conversation (newest last).')
    @action(detail=True, methods=['get'])
    def messages(self, request, pk=None):
        conversation = self.get_object()
        user = request.user
        messages = (
            conversation.messages
            .select_related('sender', 'sender__chat_key')
            .prefetch_related(
                Prefetch('reactions', queryset=MessageReaction.objects.select_related('user')),
                Prefetch(
                    'key_wraps',
                    queryset=MessageKeyWrap.objects.filter(recipient=user),
                ),
                'attachments',
                # Feeds MessageSerializer._other_participants without a per-row query.
                'conversation__participants__chat_key',
            )
            .annotate(reaction_count_anno=Count('reactions'))
            .order_by('created_at')
        )
        page = self.paginate_queryset(messages)
        ctx = {'request': request}
        if page is not None:
            return self.get_paginated_response(MessageSerializer(page, many=True, context=ctx).data)
        return Response(MessageSerializer(messages, many=True, context=ctx).data)

    @extend_schema(description='Mark all incoming messages in a conversation as read.')
    @action(detail=True, methods=['post'])
    def mark_read(self, request, pk=None):
        from django.utils import timezone
        conversation = self.get_object()
        conversation.messages.filter(is_read=False).exclude(sender=request.user).update(
            is_read=True, read_at=timezone.now(),
        )
        ConversationParticipant.objects.filter(
            conversation=conversation, user=request.user,
        ).update(last_read_at=timezone.now())
        return Response({'status': 'marked as read'})


class MessageViewSet(viewsets.ModelViewSet):
    permission_classes = [permissions.IsAuthenticated]
    pagination_class = StandardResultsSetPagination

    def get_queryset(self):
        user = self.request.user
        organization = get_active_organization(self.request)
        conversation_id = self.request.query_params.get('conversation_id')
        queryset = (
            Message.objects
            .filter(conversation__participants=user, conversation__organization=organization)
            .select_related('sender', 'sender__chat_key', 'conversation')
            .prefetch_related(
                Prefetch('reactions', queryset=MessageReaction.objects.select_related('user')),
                Prefetch('key_wraps', queryset=MessageKeyWrap.objects.filter(recipient=user)),
                'attachments',
                'conversation__participants__chat_key',
            )
            .annotate(reaction_count_anno=Count('reactions'))
        )
        if conversation_id:
            queryset = queryset.filter(conversation_id=conversation_id)
        return queryset.distinct()

    def get_serializer_class(self):
        if self.action == 'create':
            return MessageCreateSerializer
        return MessageSerializer

    def get_serializer_context(self):
        return {**super().get_serializer_context(), 'request': self.request}

    def perform_create(self, serializer):
        message = serializer.save(sender=self.request.user)
        self._broadcast_new_message(message)
        self._notify_participants(message)

    def _notify_participants(self, message):
        """In-app notification to other participants (content is E2E-encrypted, so generic text)."""
        try:
            from apps.notifications.handler import NotificationHandler
            actor = self.request.user
            actor_name = actor.get_full_name() or actor.email
            recipients = message.conversation.participants.exclude(id=actor.id)
            for u in recipients:
                NotificationHandler.send(
                    user=u,
                    title="New message",
                    message=f"{actor_name} sent you a message.",
                    notification_type="chat",
                    priority="medium",
                    icon="message-circle",
                    link=f"/chat?c={message.conversation_id}",
                    related_user=actor,
                    broadcast=True,
                )
        except Exception:
            pass

    def _broadcast_new_message(self, message):
        """No-op unless settings.WEBSOCKET_ENABLED — the WSGI-only production
        host has no consumer attached, and the thread's 5s poll is what
        actually delivers the message."""
        from django.conf import settings

        if not getattr(settings, "WEBSOCKET_ENABLED", False):
            return
        from channels.layers import get_channel_layer
        from asgiref.sync import async_to_sync
        layer = get_channel_layer()
        if not layer:
            return
        payload = MessageSerializer(message, context={'request': self.request}).data
        async_to_sync(layer.group_send)(
            f'chat_{message.conversation_id}',
            {'type': 'chat_message', 'message': payload},
        )

    def create(self, request, *args, **kwargs):
        serializer = self.get_serializer(data=request.data)
        serializer.is_valid(raise_exception=True)
        self.perform_create(serializer)
        message = serializer.instance
        read_serializer = MessageSerializer(message, context={'request': request})
        return Response(read_serializer.data, status=status.HTTP_201_CREATED)

    @extend_schema(description='Add or remove a reaction on a message.')
    @action(detail=True, methods=['post'])
    def react(self, request, pk=None):
        message = self.get_object()
        reaction = request.data.get('reaction')
        if not reaction:
            return Response({'error': 'reaction is required'}, status=status.HTTP_400_BAD_REQUEST)
        reaction_obj, created = MessageReaction.objects.get_or_create(
            message=message, user=request.user, reaction=reaction,
        )
        if not created:
            reaction_obj.delete()
            return Response({'status': 'reaction removed'})
        return Response(MessageReactionSerializer(reaction_obj).data, status=status.HTTP_201_CREATED)


class MessageAttachmentViewSet(viewsets.ModelViewSet):
    """Upload/list encrypted attachments. Server never sees plaintext."""
    permission_classes = [permissions.IsAuthenticated]
    parser_classes = [MultiPartParser, FormParser]
    http_method_names = ['get', 'post', 'delete', 'head', 'options']

    def get_queryset(self):
        return MessageAttachment.objects.filter(
            message__conversation__participants=self.request.user,
            message__conversation__organization=get_active_organization(self.request),
        ).select_related('message').distinct()

    def get_serializer_class(self):
        if self.action == 'create':
            return MessageAttachmentUploadSerializer
        return MessageAttachmentSerializer

    def perform_create(self, serializer):
        attachment = serializer.save()
        if not attachment.message.has_attachments:
            attachment.message.has_attachments = True
            attachment.message.save(update_fields=['has_attachments'])


class LinkPreviewView(APIView):
    """Fetch Open Graph metadata for one URL, requested client-side after
    decrypting a message. See link_preview.py for the SSRF hardening —
    the server has no visibility into message content, only whichever
    single URL the client asks to preview."""
    permission_classes = [permissions.IsAuthenticated]
    throttle_classes = [ScopedRateThrottle]
    throttle_scope = 'link_preview'

    @extend_schema(
        parameters=[OpenApiParameter(name='url', type=str, required=True, description='URL to fetch a preview for')],
    )
    def get(self, request):
        url = request.query_params.get('url', '').strip()
        if not url:
            return Response({'error': 'url is required'}, status=status.HTTP_400_BAD_REQUEST)
        try:
            preview = fetch_link_preview(url)
        except LinkPreviewError as exc:
            return Response({'error': exc.message}, status=exc.status_code)
        return Response(preview)
