"""Invite-only accounts: argon2id passwords, server-side sessions in an
HttpOnly cookie, and a CSRF header on unsafe requests.

Every permission check goes through `current_user` + `require_center` /
`require_admin`, so an external identity provider (e.g. NAC's Stytch) can
replace this module later without touching the routers.
"""

from __future__ import annotations

import hashlib
import secrets
import time
from collections import defaultdict
from dataclasses import dataclass, field
from datetime import UTC, datetime, timedelta

from argon2 import PasswordHasher
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
from fastapi import Depends, HTTPException, Request, Response
from psycopg import Connection

from .db import fetchall, fetchone, get_conn
from .settings import get_settings

COOKIE = "avyg_session"
CSRF_HEADER = "x-csrf-token"
SAFE_METHODS = {"GET", "HEAD", "OPTIONS"}
MIN_PASSWORD = 10

_ph = PasswordHasher()


def hash_password(password: str) -> str:
    if len(password) < MIN_PASSWORD:
        raise HTTPException(422, f"Password must be at least {MIN_PASSWORD} characters.")
    return _ph.hash(password)


def verify_password(password_hash: str, password: str) -> bool:
    try:
        return _ph.verify(password_hash, password)
    except (VerifyMismatchError, VerificationError, InvalidHashError):
        return False


def token_hash(token: str) -> str:
    return hashlib.sha256(token.encode()).hexdigest()


# --- login throttling (per process; fine for a handful of users) -------------

_ATTEMPTS: dict[str, list[float]] = defaultdict(list)
WINDOW_S = 15 * 60
MAX_ATTEMPTS = 10


def check_throttle(key: str) -> None:
    now = time.monotonic()
    _ATTEMPTS[key] = [t for t in _ATTEMPTS[key] if now - t < WINDOW_S]
    if len(_ATTEMPTS[key]) >= MAX_ATTEMPTS:
        raise HTTPException(429, "Too many login attempts. Try again in 15 minutes.")


def record_failure(key: str) -> None:
    _ATTEMPTS[key].append(time.monotonic())


def clear_throttle(key: str) -> None:
    _ATTEMPTS.pop(key, None)


# --- sessions ------------------------------------------------------------------


def create_session(conn: Connection, user_id: int, response: Response) -> str:
    settings = get_settings()
    token = secrets.token_urlsafe(32)
    csrf = secrets.token_urlsafe(24)
    expires = datetime.now(UTC) + timedelta(days=settings.session_days)
    conn.execute(
        "INSERT INTO sessions (id, user_id, csrf_token, expires_at) VALUES (%s, %s, %s, %s)",
        (token_hash(token), user_id, csrf, expires),
    )
    conn.execute("UPDATE users SET last_login_at = now() WHERE id = %s", (user_id,))
    conn.execute("DELETE FROM sessions WHERE expires_at < now()")
    response.set_cookie(
        COOKIE,
        token,
        max_age=settings.session_days * 86400,
        httponly=True,
        secure=settings.cookie_secure,
        samesite="lax",
        path="/",
    )
    return csrf


def destroy_session(conn: Connection, request: Request, response: Response) -> None:
    token = request.cookies.get(COOKIE)
    if token:
        conn.execute("DELETE FROM sessions WHERE id = %s", (token_hash(token),))
    response.delete_cookie(COOKIE, path="/")


@dataclass
class User:
    id: int
    email: str
    name: str
    is_admin: bool
    csrf: str
    memberships: dict[int, str] = field(default_factory=dict)

    def role_in(self, center_id: int) -> str | None:
        if self.is_admin:
            return "center_admin"
        return self.memberships.get(center_id)


def load_memberships(conn: Connection, user_id: int) -> dict[int, str]:
    rows = fetchall(conn, "SELECT center_id, role FROM memberships WHERE user_id = %s", (user_id,))
    return {r["center_id"]: r["role"] for r in rows}


def current_user(request: Request, conn: Connection = Depends(get_conn)) -> User:
    token = request.cookies.get(COOKIE)
    if not token:
        raise HTTPException(401, "Not signed in")
    row = fetchone(
        conn,
        """SELECT s.id AS sid, s.csrf_token, s.last_seen_at, u.id, u.email, u.name, u.is_admin
           FROM sessions s JOIN users u ON u.id = s.user_id
           WHERE s.id = %s AND s.expires_at > now() AND u.is_active""",
        (token_hash(token),),
    )
    if not row:
        raise HTTPException(401, "Session expired")
    if request.method not in SAFE_METHODS and not secrets.compare_digest(
        request.headers.get(CSRF_HEADER, ""), row["csrf_token"]
    ):
        raise HTTPException(403, "Missing or invalid CSRF token")
    if datetime.now(UTC) - row["last_seen_at"] > timedelta(minutes=10):
        conn.execute("UPDATE sessions SET last_seen_at = now() WHERE id = %s", (row["sid"],))
    return User(
        id=row["id"],
        email=row["email"],
        name=row["name"],
        is_admin=row["is_admin"],
        csrf=row["csrf_token"],
        memberships=load_memberships(conn, row["id"]),
    )


def require_center(user: User, center_id: int, admin: bool = False) -> str:
    """Raise 403 unless the user belongs to the center (as center_admin if admin=True)."""
    role = user.role_in(center_id)
    if role is None or (admin and role != "center_admin"):
        raise HTTPException(403, "You do not have access to this center")
    return role


def require_admin(user: User) -> None:
    if not user.is_admin:
        raise HTTPException(403, "Administrators only")


def me_payload(conn: Connection, user: User) -> dict:
    centers = fetchall(
        conn,
        "SELECT id, slug, name FROM centers ORDER BY name"
        if user.is_admin
        else "SELECT c.id, c.slug, c.name FROM centers c JOIN memberships m ON m.center_id = c.id WHERE m.user_id = %s ORDER BY c.name",
        None if user.is_admin else (user.id,),
    )
    return {
        "user": {"id": user.id, "email": user.email, "name": user.name, "is_admin": user.is_admin},
        "centers": [{**c, "role": user.role_in(c["id"])} for c in centers],
        "csrf": user.csrf,
    }
