import os
import jwt
import requests

from django.db import transaction
from django.utils import timezone
from datetime import timedelta

from rest_framework import status
from rest_framework.exceptions import ValidationError

from api.models.access_token import AccessToken
from api.models.user import User
from api.views.admin.base import AdminApiBase
from api.validator import BaseValidator
from api.translation import __
from api.utils import log_critical_error, get_exception_detail


MICROSOFT_KEYS_URL = "https://login.microsoftonline.com/common/discovery/v2.0/keys"

ENTRA_CLIENT_ID = os.getenv("ENTRA_CLIENT_ID", "")
ENTRA_TENANT_ID = os.getenv("ENTRA_TENANT_ID", "")
ENTRA_DEFAULT_USER_TYPE = os.getenv("ENTRA_DEFAULT_USER_TYPE", "admin")


class EntraLogin(AdminApiBase):
    authentication_classes = []

    def validate_request(self, request):
        rules = {
            "id_token": "required",
        }

        attribute = {}

        validator = BaseValidator(request.data, rules, attribute)
        error = validator.validate()
        if error:
            raise ValidationError(error)

    def raise_login_failed(self):
        raise ValidationError({
            "error": __("validation.login_failed"),
            "errors": {
                "id_token": __("validation.login_failed")
            }
        })

    def get_jwt_header(self, id_token):
        try:
            return jwt.get_unverified_header(id_token)
        except Exception:
            self.raise_login_failed()

    def fetch_microsoft_keys(self):
        response = requests.get(MICROSOFT_KEYS_URL, timeout=10)
        response.raise_for_status()
        return response.json().get("keys", [])

    def find_public_key_by_kid(self, keys, kid):
        for key in keys:
            if key.get("kid") == kid:
                return key
        return None

    def decode_microsoft_token(self, id_token):
        try:
            header = self.get_jwt_header(id_token)
            kid = header.get("kid")

            keys = self.fetch_microsoft_keys()
            jwk = self.find_public_key_by_kid(keys, kid)

            if not jwk:
                raise Exception("Microsoft public key not found")

            public_key = jwt.algorithms.RSAAlgorithm.from_jwk(jwk)

            decode_options = {}

            if not ENTRA_CLIENT_ID:
                decode_options["verify_aud"] = False

            decoded = jwt.decode(
                id_token,
                public_key,
                algorithms=["RS256"],
                audience=ENTRA_CLIENT_ID if ENTRA_CLIENT_ID else None,
                options=decode_options,
            )

            return decoded

        except ValidationError:
            raise
        except Exception:
            self.raise_login_failed()

    def validate_tenant_id(self, tenant_id):
        # Meeting note:
        # get tenant id from id_token, compare with env one.
        # if same, login/register.
        # if not same, return user not found / login failed style.
        if not ENTRA_TENANT_ID:
            self.raise_login_failed()

        if not tenant_id:
            self.raise_login_failed()

        if tenant_id != ENTRA_TENANT_ID:
            self.raise_login_failed()

    def post(self, request):
        try:
            self.validate_request(request)

            id_token = request.data.get("id_token")

            claims = self.decode_microsoft_token(id_token)

            email = claims.get("email")
            name = claims.get("name") or ""
            microsoft_oid = claims.get("oid")
            tenant_id = claims.get("tid")

            self.validate_tenant_id(tenant_id)

            if not email:
                self.raise_login_failed()

            if not microsoft_oid:
                self.raise_login_failed()

            app_type = request.headers.get("X-App-Type", "web")
            device = request.headers.get("User-Agent")

            with transaction.atomic():
                user = User.objects.filter(microsoft_oid=microsoft_oid).first()

                if not user:
                    user = User.objects.filter(email=email).first()

                    if user:
                        user.microsoft_oid = microsoft_oid

                        update_fields = ["microsoft_oid", "updated_at"]

                        if not user.name and name:
                            user.name = name
                            update_fields.append("name")

                        user.save(update_fields=update_fields)

                if not user:
                    default_user_type = User.USER_TYPE.get(
                        ENTRA_DEFAULT_USER_TYPE,
                        User.USER_TYPE["admin"]
                    )

                    user = User(
                        email=email,
                        name=name,
                        microsoft_oid=microsoft_oid,
                        user_type=default_user_type,
                        status=User.STATUS_TO_CODE["active"],
                        is_only_self=True,
                        can_login_with_password=False,
                    )
                    user.set_unusable_password()
                    user.save()

                if user.user_type != User.USER_TYPE.get("admin"):
                    self.raise_login_failed()

                AccessToken.objects.filter(
                    user=user,
                    type=app_type,
                    device=device
                ).delete()

                token = AccessToken.objects.create(
                    user=user,
                    type=app_type,
                    device=device,
                    expires_at=timezone.now() + timedelta(hours=2)
                )

            response = {
                "user": {
                    "id": user.id,
                    "email": user.email,
                    "name": user.name,
                },
                "token": token.token
            }

            return self.api_response(data=response)

        except ValidationError as e:
            first_message = e.detail["error"]
            errors = e.detail["errors"]
            return self.api_response(
                error=first_message,
                errors=errors,
                code=status.HTTP_400_BAD_REQUEST
            )

        except Exception as e:
            e_details = get_exception_detail(e)
            log_critical_error(
                user_id=None,
                descr=e_details["descr"],
                url=e_details["url"],
                trace=e_details["trace"]
            )
            return self.api_response(
                error=__("message.internal_server_error"),
                code=status.HTTP_500_INTERNAL_SERVER_ERROR
            )