"""Authentification : connexion, MFA, sessions, limitation des tentatives."""
from __future__ import annotations

import time
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone

from sqlalchemy import select, update
from sqlalchemy.orm import Session

from ..core import audit, security
from ..models import AuthThrottle, RecoveryCode, RolePermission, User, UserSession
from . import settings as settings_service

GENERIC_LOGIN_ERROR = "Identifiant ou mot de passe incorrect."
LOCKED_ERROR = "Trop de tentatives. Réessayez dans quelques minutes."
IP_MAX_FAILURES = 30


def utcnow() -> datetime:
    return datetime.now(timezone.utc)


@dataclass
class Principal:
    user: User
    session: UserSession
    permissions: frozenset[str]
    mfa_required: bool

    @property
    def needs_mfa_setup(self) -> bool:
        return self.mfa_required and not self.user.mfa_enabled

    @property
    def needs_password_change(self) -> bool:
        return self.user.must_change_password

    @property
    def fully_authenticated(self) -> bool:
        return self.session.mfa_verified and not self.needs_mfa_setup and not self.needs_password_change

    def can(self, perm: str) -> bool:
        return perm in self.permissions


def permissions_for(db: Session, user: User) -> frozenset[str]:
    perms = set(db.execute(
        select(RolePermission.permission_code).where(RolePermission.role_code == user.role_code)
    ).scalars())
    if user.role_code == "reader" and user.can_export:
        perms.add("export.create")
    return frozenset(perms)


def mfa_required_for(db: Session, user: User) -> bool:
    key = "mfa_required_admin" if user.role_code == "admin" else "mfa_required_reader"
    return bool(settings_service.get(db, key))


# --- Limitation des tentatives -------------------------------------------------------------------

def _throttle_keys(username: str, ip: str | None) -> list[tuple[str, int]]:
    keys = [(f"u:{username.lower()[:120]}", 0)]
    if ip:
        keys.append((f"ip:{ip}", IP_MAX_FAILURES))
    return keys


def is_locked(db: Session, username: str, ip: str | None) -> bool:
    now = utcnow()
    for key, _ in _throttle_keys(username, ip):
        row = db.get(AuthThrottle, key)
        if row is not None and row.locked_until is not None and row.locked_until > now:
            return True
    return False


def register_failure(db: Session, username: str, ip: str | None) -> int:
    """Incrémente les compteurs ; verrouille au-delà du seuil. Renvoie le nombre d'échecs du compte."""
    now = utcnow()
    max_failures = int(settings_service.get(db, "login_max_failures"))
    lock_minutes = int(settings_service.get(db, "login_lockout_minutes"))
    user_failures = 0
    for key, limit in _throttle_keys(username, ip):
        limit = limit or max_failures
        row = db.get(AuthThrottle, key, with_for_update=True)
        if row is None:
            row = AuthThrottle(key=key, failures=0, first_failure_at=now, lock_count=0)
            db.add(row)
        elif now - row.first_failure_at > timedelta(hours=1) and not (row.locked_until and row.locked_until > now):
            row.failures, row.first_failure_at = 0, now
        row.failures += 1
        if row.failures >= limit:
            # Verrouillage progressif : 1×, 2×, 4× la durée de base (plafonné à 24 h).
            factor = 2 ** min(row.lock_count, 6)
            row.locked_until = now + timedelta(minutes=min(lock_minutes * factor, 1440))
            row.lock_count += 1
            row.failures = 0
            row.first_failure_at = now
        if key.startswith("u:"):
            user_failures = row.failures
    db.flush()
    return user_failures


def clear_failures(db: Session, username: str) -> None:
    row = db.get(AuthThrottle, f"u:{username.lower()[:120]}")
    if row is not None:
        db.delete(row)


def progressive_delay(failures: int) -> None:
    if failures > 1:
        time.sleep(min(0.25 * 2 ** (failures - 2), 3.0))


# --- Sessions ------------------------------------------------------------------------------------

def create_session(db: Session, user: User, *, mfa_verified: bool, ip: str | None, user_agent: str | None
                   ) -> tuple[str, UserSession]:
    token = security.new_token()
    hours = int(settings_service.get(db, "session_absolute_hours"))
    sess = UserSession(
        id=security.token_digest(token), user_id=user.id, csrf_token=security.new_token(),
        mfa_verified=mfa_verified, expires_at=utcnow() + timedelta(hours=hours),
        ip=ip, user_agent=(user_agent or "")[:200], last_seen_at=utcnow(),
    )
    db.add(sess)
    db.flush()
    return token, sess


def rotate_session(db: Session, old: UserSession, user: User, *, mfa_verified: bool) -> tuple[str, UserSession]:
    token, sess = create_session(db, user, mfa_verified=mfa_verified, ip=old.ip, user_agent=old.user_agent)
    sess.expires_at = old.expires_at  # la rotation ne prolonge pas la durée absolue
    old.revoked_at, old.revoked_reason = utcnow(), "rotation"
    return token, sess


def load_session(db: Session, token: str | None) -> tuple[UserSession, User] | None:
    if not token or len(token) > 100:
        return None
    sess = db.get(UserSession, security.token_digest(token))
    if sess is None or sess.revoked_at is not None:
        return None
    now = utcnow()
    idle = int(settings_service.get(db, "session_idle_minutes"))
    if sess.expires_at <= now or sess.last_seen_at + timedelta(minutes=idle) <= now:
        sess.revoked_at, sess.revoked_reason = now, "expiration"
        db.commit()
        return None
    user = db.get(User, sess.user_id)
    if user is None or not user.is_active:
        return None
    if now - sess.last_seen_at > timedelta(seconds=30):
        sess.last_seen_at = now
        db.commit()
    return sess, user


def revoke_user_sessions(db: Session, user_id: int, reason: str, except_id: str | None = None) -> int:
    stmt = (update(UserSession)
            .where(UserSession.user_id == user_id, UserSession.revoked_at.is_(None))
            .values(revoked_at=utcnow(), revoked_reason=reason))
    if except_id:
        stmt = stmt.where(UserSession.id != except_id)
    return db.execute(stmt).rowcount or 0


# --- Connexion -----------------------------------------------------------------------------------

class AuthError(Exception):
    def __init__(self, message: str, status: int = 401, code: str = "auth_failed"):
        super().__init__(message)
        self.message, self.status, self.code = message, status, code


def authenticate(db: Session, username: str, password: str, ip: str | None) -> User:
    username = (username or "").strip().lower()
    if is_locked(db, username, ip):
        audit.record(db, "auth.login_blocked", username=username[:64], ip=ip)
        db.commit()
        raise AuthError(LOCKED_ERROR, 429, "locked")
    user = db.execute(select(User).where(User.username == username)).scalar_one_or_none() if username else None
    ok = security.verify_password(user.password_hash if user and user.is_active else None, password or "")
    if not ok or user is None:
        failures = register_failure(db, username, ip)
        audit.record(db, "auth.login_failed", user_id=user.id if user else None, username=username[:64], ip=ip,
                     details={"motif": "compte désactivé" if user and not user.is_active else "identifiants"})
        db.commit()
        progressive_delay(failures)
        raise AuthError(GENERIC_LOGIN_ERROR)
    clear_failures(db, username)
    if security.needs_rehash(user.password_hash):
        user.password_hash = security.hash_password(password)
    return user


def verify_second_factor(db: Session, user: User, code: str, ip: str | None) -> str:
    """Vérifie un code TOTP ou un code de récupération. Renvoie la méthode utilisée."""
    if is_locked(db, user.username, ip):
        raise AuthError(LOCKED_ERROR, 429, "locked")
    code = (code or "").strip()
    if user.mfa_enabled and user.mfa_secret_enc:
        step = security.verify_totp(security.decrypt_secret(user.mfa_secret_enc), code, user.mfa_last_step)
        if step is not None:
            user.mfa_last_step = step
            clear_failures(db, user.username)
            return "totp"
        digest = security.recovery_digest(code)
        rc = db.execute(select(RecoveryCode).where(
            RecoveryCode.user_id == user.id, RecoveryCode.code_hash == digest, RecoveryCode.used_at.is_(None)
        )).scalar_one_or_none()
        if rc is not None and len(code) >= 12:
            rc.used_at = utcnow()
            clear_failures(db, user.username)
            return "recovery"
    failures = register_failure(db, user.username, ip)
    audit.record(db, "auth.mfa_failed", user_id=user.id, username=user.username, ip=ip)
    db.commit()
    progressive_delay(failures)
    raise AuthError("Code incorrect ou expiré.", 401, "mfa_failed")


def issue_recovery_codes(db: Session, user: User) -> list[str]:
    for rc in db.execute(select(RecoveryCode).where(RecoveryCode.user_id == user.id)).scalars():
        db.delete(rc)
    codes = security.new_recovery_codes()
    for c in codes:
        db.add(RecoveryCode(user_id=user.id, code_hash=security.recovery_digest(c)))
    return codes


def remaining_recovery_codes(db: Session, user_id: int) -> int:
    return len(db.execute(select(RecoveryCode.id).where(
        RecoveryCode.user_id == user_id, RecoveryCode.used_at.is_(None))).all())
