from datetime import timedelta
import hashlib
import time
from http import HTTPStatus

from django.contrib.contenttypes.models import ContentType
from django.core.files.uploadedfile import UploadedFile
from django.db import models
from rest_framework.views import APIView
from rest_framework.authentication import BaseAuthentication
from rest_framework.exceptions import AuthenticationFailed,ValidationError,PermissionDenied
from rest_framework.permissions import IsAuthenticated
from django.utils import timezone
from api.models.access_token import AccessToken
from rest_framework.response import Response

from api.models.incoming_api_admin import IncomingApiAdmin
from api.utils import get_ip,throw_validation_error
from django.conf import settings
from api.translation import __

class TokenAuthentication(BaseAuthentication):
    def authenticate(self, request):
        auth_header = request.headers.get('Authorization')
        if not auth_header or not auth_header.startswith("Token "):
            raise AuthenticationFailed(__("validation.invalid",field=__("attr.token")))

        token_key = auth_header.split(" ")[1]
        try:
            token = AccessToken.objects.select_related('user').get(token=token_key)
        except AccessToken.DoesNotExist:
            raise AuthenticationFailed(__("validation.invalid",field=__("attr.token")))

        if token.expires_at and token.expires_at <= timezone.now():
            token.delete()
            raise AuthenticationFailed(__("validation.token_expired"))

        return (token.user, token)

class AdminApiBase(APIView):
    authentication_classes = [TokenAuthentication]

    def finalize_response(self, request, response, *args, **kwargs):
        if 200 <= response.status_code < 300:
            if not self.exclude_token_extend:
                token = getattr(request, "auth", None)
                if isinstance(token, AccessToken):
                    token.expires_at = timezone.now() + timedelta(hours=1)
                    token.save(update_fields=["expires_at"])

        return super().finalize_response(request, response, *args, **kwargs)
    
    enable_log = True
    exclude_token_extend = False
    kafka_config = settings.KAFKA_CONFIG
    # if want put desc, "-name" will handle desc order, if "name" will be asc order
    list_default_order_column = "name"
    # for listing log, no need store the outgoing data of the list
    api_log_skip_outgoing_data = False
    # audit_type=AuditTrail.TYPE['admin']
    url_name_ignore_basic_validation_and_enable_log = [
        "get_timestamp"
    ]

    def dispatch(self, request, *args, **kwargs):
        return super().dispatch(request, *args, **kwargs)

    def initial(self, request, *args, **kwargs):
        super().initial(request, *args, **kwargs)
        self.admin_api_log_id = None
        url_name = request.resolver_match.view_name or ""
        if url_name not in self.url_name_ignore_basic_validation_and_enable_log:
            if getattr(self, 'enable_log', True):
                # no need permission checking for calendar first
                # route_name = request.resolver_match.view_name
                #
                # if route_name not in NO_PERMISSION:
                #     if not request.user.has_perm(route_name):
                #         raise PermissionDenied(__('validation.no_permission'))

                incoming_data = request.data.copy()
                remove_fields = ['password']
                for field in remove_fields:
                    incoming_data.pop(field, None)

                self.admin_api_log_id = IncomingApiAdmin.insert_log(
                    user_id=getattr(request.user, "id", None),
                    request_data=incoming_data,
                    url=request.scheme + "://" + request.get_host() + request.path,
                    ip=get_ip(request),
                )

            #some general validation
            errors = self.basic_validation(request.data)
            if errors:
                scode = HTTPStatus.BAD_REQUEST
                IncomingApiAdmin.update_log(self.admin_api_log_id, errors, scode=scode)
                raise throw_validation_error(code=scode, error=errors)

    def basic_validation(self, data: dict) -> bool:
        if "timestamp" not in data:
            return __("validation.required", field=__("attr.timestamp"))

        if "signature" not in data:
            return __("validation.required", field=__("attr.signature"))

        #validate timestamp
        request_timestamp = int(data["timestamp"])
        now = int(time.time())
        min_time = now - settings.API_TIMEOUT_MAX
        max_time = now - settings.API_TIMEOUT_MIN
        if request_timestamp > max_time or request_timestamp < min_time:
            return __("validation.invalid",field=__("attr.timestamp"))

        #validate signature
        request_sign = data.get('signature')
        self_sign = self.sign(data, settings.API_SECRET_KEY)
        if request_sign != self_sign:
            return __("validation.invalid",field=__("signature"))

    def api_response(self, data=None, code=200, errors=None, error=None):
        response = {
            'code': code,
        }
        if errors or error:
            response['error'] = error
            response['errors'] = errors
        else:
            data = data.copy()

            remove_fields = ['password']
            for field in remove_fields:
                data.pop(field, None)

            response['data'] = data

        if getattr(self, 'enable_log', True) and self.admin_api_log_id:
            outgoing_data = response.copy()
            if self.api_log_skip_outgoing_data:
                outgoing_data = None
            IncomingApiAdmin.update_log(self.admin_api_log_id, outgoing_data, scode=code)

        return Response(response, status=code)

    @staticmethod
    def sign(data: dict, secret: str) -> str:
        # Copy data to avoid modifying the original
        sign_data = data

        # Remove 'sign' and 'signature' keys if present
        sign_data.pop('sign', None)
        sign_data.pop('signature', None)
        # Sort by key
        sorted_items = sorted(sign_data.items())

        # Concatenate non-list/dict values
        raw_sign = ''
        for key, val in sorted_items:
            # null,None,"",false,0 also skip
            if val:
                if isinstance(val, (list, dict, UploadedFile)):
                    continue  # skip arrays and upload file
                raw_sign += str(val)

        # Append secret
        raw_sign += secret

        # Return SHA256 hex digest
        return hashlib.sha256(raw_sign.encode('utf-8')).hexdigest()

    @staticmethod
    def update_m2m_field(model, field_name, old_ids, new_ids, old_new_data, two_side=False, generic=False):
        """
        :param edit_model: model want to edit
        :param field_name: the field name eg.suitability
        :param old_ids: list of old ids
        :param new_ids: list of new ids
        :param old_new_data: dict of old_new_data
        :param two_side: when insert/delete a and b, will insert/delete b and a also, mostly used in same table relation liek avoid_concurrency
        :param generic: when relation only have 1 foreign key, and using refer_table and refer_id, eg. tag_relation table
        :return:
        """
        # if new ids pass in by FE is None, make it as empty list, so it can be compare
        if new_ids is None:
            new_ids = []
        if old_ids is None:
            old_ids = []

        to_add = [i for i in new_ids if i not in old_ids]
        to_delete = [i for i in old_ids if i not in new_ids]

        if not (to_add or to_delete):
            return old_new_data

        m2m_field = getattr(model, field_name)
        model_id = model.id
        # only have old_ids, only need old_data, if not it was a insert action, dun have old data
        if old_ids:
            old_new_data["old_data"][field_name] = list(old_ids)

        old_new_data["new_data"][field_name] = list(new_ids)
        if generic:
            # Generic relation handling
            ct = ContentType.objects.get_for_model(model)
            # DELETE
            if to_delete:
                model_class = getattr(model, field_name).model
                model_class.objects.filter(
                    refer_table=ct,
                    refer_id=model.id,
                    **{f"{field_name}_id__in": to_delete}
                ).delete()

            # INSERT
            if to_add:
                model_class = getattr(model, field_name).model
                objs = [
                    model_class(refer_table=ct, refer_id=model.id, **{f"{field_name}_id": pk})
                    for pk in to_add
                ]
                model_class.objects.bulk_create(objs, ignore_conflicts=True)
        elif two_side:
            through = m2m_field.through
            src_field = m2m_field.source_field_name
            tgt_field = m2m_field.target_field_name
            # DELETE both directions (2 queries)
            if to_delete:
                through.objects.filter(**{src_field: model, f"{tgt_field}__in": to_delete}).delete()
                through.objects.filter(**{tgt_field: model, f"{src_field}__in": to_delete}).delete()
            # INSERT both directions (bulk_create)
            if to_add:
                # Fetch instances once
                target_model = m2m_field.model
                target_instances = {s.id: s for s in target_model.objects.filter(id__in=to_add)}
                two_side_objs = []
                for pk in to_add:
                    other_instance = target_instances[pk]
                    two_side_objs.append(through(**{src_field: model, tgt_field: other_instance}))
                    two_side_objs.append(through(**{src_field: other_instance, tgt_field: model}))
                through.objects.bulk_create(two_side_objs, ignore_conflicts=True)
        else:
            # normal relation update
            if to_delete:
                m2m_field.remove(*to_delete)
            if to_add:
                m2m_field.add(*to_add)

        return old_new_data

    @staticmethod
    def update_m2m_field_bulk(models, field_name, new_ids, old_new_data, two_side=False, generic=False):
        """
        Bulk update ManyToMany relation for one or multiple models, optimized for prefetched data.

        :param models: single model or queryset/list of models (with prefetch_related(field_name))
        :param field_name: the field name eg. 'suitability'
        :param new_ids: list of new ids (same for all models)
        :param old_new_data: dict of old_new_data
        :param two_side: when relation is symmetric (eg. avoid_concurrency)
        :param generic: when relation table uses refer_table + refer_id
        :return: old_new_data (with audit logs)
        """
        if not isinstance(models, (list, tuple)) and not hasattr(models, "__iter__"):
            models = [models]

        if new_ids is None:
            new_ids = []
        new_ids = list(set(new_ids))  # deduplicate

        if not models:
            return old_new_data

        model_class = models[0].__class__
        m2m_field = getattr(model_class, field_name)

        for model in models:
            relation = getattr(model, field_name)  # e.g. model.tag
            if generic:
                sub_field = field_name + "_id"
                old_ids = [getattr(obj, sub_field) for obj in relation.all()]
            else:
                # normal m2m
                old_ids = [obj.id for obj in relation.all()]
            if set(old_ids) != set(new_ids):
                old_new_data["old_data"].setdefault(model.id, {})[field_name] = old_ids
                old_new_data["new_data"].setdefault(model.id, {})[field_name] = new_ids
        # Exit early if no changes at all
        if not old_new_data.get("new_data"):
            return old_new_data

        model_ids = [model.id for model in models]

        if generic:
            # take 1 of the model in models to get content type and relation model, cause all will in same model
            model=models[0]
            ct = ContentType.objects.get_for_model(model)
            relation_model = getattr(model, field_name).model

            # DELETE outdated
            relation_model.objects.filter(refer_table=ct, refer_id__in=model_ids).exclude(
                **{f"{field_name}_id__in": new_ids}
            ).delete()

            # INSERT missing
            objs = [
                relation_model(refer_table=ct, refer_id=mid, **{f"{field_name}_id": pk})
                for mid in model_ids for pk in new_ids
            ]
            relation_model.objects.bulk_create(objs, ignore_conflicts=True)

        elif two_side:
            through = m2m_field.through
            # just want to get src_field and target field
            get_field_m2m = getattr(models[0], field_name)
            src_field = f"{get_field_m2m.source_field_name}_id"
            tgt_field = f"{get_field_m2m.target_field_name}_id"

            # DELETE both directions
            through.objects.filter(**{f"{src_field}__in": model_ids}).exclude(
                **{f"{tgt_field}__in": new_ids}
            ).delete()
            through.objects.filter(**{f"{tgt_field}__in": model_ids}).exclude(
                **{f"{src_field}__in": new_ids}
            ).delete()

            # INSERT both directions
            objs = []
            for mid in model_ids:
                for pk in new_ids:
                    objs.append(through(**{src_field: mid, tgt_field: pk}))
                    objs.append(through(**{src_field: pk, tgt_field: mid}))
            through.objects.bulk_create(objs, ignore_conflicts=True)

        else:
            # Normal M2M relation
            through = m2m_field.through
            # just want to get src_field and target field
            get_field_m2m = getattr(models[0], field_name)
            src_field = f"{get_field_m2m.source_field_name}_id"
            tgt_field = f"{get_field_m2m.target_field_name}_id"

            # DELETE outdated
            through.objects.filter(**{f"{src_field}__in": model_ids}).exclude(
                **{f"{tgt_field}__in": new_ids}
            ).delete()

            # INSERT missing
            objs = [through(**{src_field: mid, tgt_field: pk}) for mid in model_ids for pk in new_ids]
            through.objects.bulk_create(objs, ignore_conflicts=True)

        return old_new_data