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.audit_trail import AuditTrail
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 __
from api.admin_permission_conf import PERMISSIONS,NO_PERMISSION

class TokenAuthentication(BaseAuthentication):
    def authenticate(self, request):
        auth_header = request.headers.get('Authorization')
        if not auth_header or not auth_header.startswith("Token "):
            return None

        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"))

        # Auto-extend session on each API call (sliding window, same as login expiry)
        token.extend_expiry(minutes=120)

        return (token.user, token)

class AdminApiBase(APIView):
    authentication_classes = [TokenAuthentication]
    enable_log = True
    kafka_config = settings.KAFKA_CONFIG
    default_per_page = -1
    default_page = 1
    # 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']
    default_code_prefix = settings.DEFAULT_CODE_PREFIX
    # when the code or name is auto generate, will used this param to seperate, eg. prefix is m21 so will be m21/
    auto_generate_name_seperator = "/"

    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
        if getattr(self, 'enable_log', True):
            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
        
    def parse_pagination(self, request):
        page = int(request.data.get("page") or self.default_page)
        per_page = int(request.data.get("per_page") or self.default_per_page)
        return page, per_page

    def listing_pagination(self,request,model):
        page, per_page = self.parse_pagination(request)
        sort_by = request.data.get("sort_by")
        order_by = request.data.get("order_by")
        total = model.count()

        if sort_by:
            clean_default = self.list_default_order_column.lstrip("-")
            if str(order_by).lower() == "desc":
                ordering_string = f"-{sort_by}"
                secondary_string = f"-{clean_default}"
            else:
                ordering_string = sort_by
                secondary_string = clean_default

            if sort_by != clean_default:
                model = model.order_by(ordering_string, secondary_string)
            else:
                model = model.order_by(ordering_string)
        else:
            model = model.order_by(self.list_default_order_column)

        # -1 means all
        if per_page != -1:
            start_index = (page - 1) * per_page
            end_index = start_index + per_page
            paginated_data = model[start_index:end_index]
        else:
            paginated_data = model

        return {
            "total": total,
            "paginated_data":paginated_data,
            "page": page,
            "per_page": per_page,
        }

    def listing_filter(self,request,model):
        filters = request.data.get("filter")
        if filters:
            for key,val in filters.items():
                match key:
                    case "code"|"name": # using LIKE %%
                        filter_kwargs = {f"{key}__icontains": val}
                        model = model.filter(**filter_kwargs)
                    case "department_id": # must be array only can do
                        if isinstance(val, (list, tuple)):
                            model = model.filter(department_id__in=val)
                        else:
                            filter_kwargs = {key: val}
                            model = model.filter(**filter_kwargs)
        return model