from __future__ import annotations

import logging
import random
import secrets
import smtplib
import time
from email.message import EmailMessage
from threading import Lock

from app.config import settings

logger = logging.getLogger(__name__)

OTP_TTL_SECONDS = 10 * 60
OTP_LENGTH = 6
MAX_ATTEMPTS = 5


class OtpService:
    def __init__(self) -> None:
        self._lock = Lock()
        # key: f"{platform}:{email}" -> record
        self._otps: dict[str, dict] = {}
        self._reset_tokens: dict[str, dict] = {}

    @staticmethod
    def _key(platform: str, email: str) -> str:
        return f"{platform}:{email.strip().lower()}"

    def create_otp(self, platform: str, email: str) -> str:
        code = "".join(str(random.randint(0, 9)) for _ in range(OTP_LENGTH))
        key = self._key(platform, email)
        with self._lock:
            self._otps[key] = {
                "code": code,
                "expires_at": time.time() + OTP_TTL_SECONDS,
                "attempts": 0,
            }
        return code

    def verify_otp(self, platform: str, email: str, otp: str) -> str | None:
        key = self._key(platform, email)
        with self._lock:
            record = self._otps.get(key)
            if not record:
                return None
            if time.time() > record["expires_at"]:
                self._otps.pop(key, None)
                return None
            record["attempts"] += 1
            if record["attempts"] > MAX_ATTEMPTS:
                self._otps.pop(key, None)
                return None
            if record["code"] != otp.strip():
                return None
            self._otps.pop(key, None)
            token = secrets.token_urlsafe(24)
            self._reset_tokens[token] = {
                "platform": platform,
                "email": email.strip().lower(),
                "expires_at": time.time() + OTP_TTL_SECONDS,
            }
            return token

    def consume_reset_token(self, token: str, platform: str, email: str) -> bool:
        with self._lock:
            record = self._reset_tokens.get(token)
            if not record:
                return False
            if time.time() > record["expires_at"]:
                self._reset_tokens.pop(token, None)
                return False
            if record["platform"] != platform or record["email"] != email.strip().lower():
                return False
            self._reset_tokens.pop(token, None)
            return True


otp_service = OtpService()


def send_otp_email(to_email: str, otp: str, platform: str) -> bool:
    """Send OTP email when SMTP is configured. Returns True if sent."""
    if not settings.smtp_host or not settings.smtp_from:
        logger.info(
            "SMTP not configured — OTP for %s (%s) is %s (prototype mode)",
            to_email,
            platform,
            otp,
        )
        return False

    subject = "Your Amanah password reset code"
    body = (
        f"Your one-time password (OTP) for Amanah ({platform}) is:\n\n"
        f"    {otp}\n\n"
        f"This code expires in {OTP_TTL_SECONDS // 60} minutes.\n"
        "If you did not request this, you can ignore this email.\n"
    )

    message = EmailMessage()
    message["Subject"] = subject
    message["From"] = settings.smtp_from
    message["To"] = to_email
    message.set_content(body)

    try:
        with smtplib.SMTP(settings.smtp_host, settings.smtp_port, timeout=15) as smtp:
            if settings.smtp_use_tls:
                smtp.starttls()
            if settings.smtp_user and settings.smtp_password:
                smtp.login(settings.smtp_user, settings.smtp_password)
            smtp.send_message(message)
        logger.info("OTP email sent to %s", to_email)
        return True
    except Exception:
        logger.exception("Failed to send OTP email to %s", to_email)
        return False
