"""Small, testable authentication core for the OpenTask auth-flow sample.

The sample intentionally keeps storage in memory.  Replace the two stores with
durable, encrypted-at-rest storage in a real deployment and rotate the JWT key
outside source control.
"""

from __future__ import annotations

import base64
import hashlib
import hmac
import json
import re
import secrets
import time
from dataclasses import dataclass
from typing import Any

import bcrypt


USERNAME_RE = re.compile(r"^[A-Za-z0-9_.-]{3,64}$")
MIN_PASSWORD_LENGTH = 12


class AuthError(ValueError):
    """A safe client-facing authentication error."""


@dataclass(frozen=True)
class Session:
    user: str
    csrf: str
    expires_at: float


def _b64(value: bytes) -> str:
    return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii")


def _unb64(value: str) -> bytes:
    return base64.urlsafe_b64decode(value + "=" * ((4 - len(value) % 4) % 4))


class AuthService:
    """In-memory registration/login service with opaque sessions and CSRF."""

    def __init__(self, *, jwt_key: bytes | None = None, session_ttl: int = 3600) -> None:
        if session_ttl < 1:
            raise ValueError("session_ttl must be positive")
        self._passwords: dict[str, bytes] = {}
        self._sessions: dict[str, Session] = {}
        self._jwt_key = jwt_key or secrets.token_bytes(32)
        self._session_ttl = session_ttl

    @staticmethod
    def _validate_user(username: str) -> str:
        if not isinstance(username, str) or not USERNAME_RE.fullmatch(username):
            raise AuthError("invalid username")
        return username.casefold()

    @staticmethod
    def _validate_password(password: str) -> str:
        if not isinstance(password, str) or len(password) < MIN_PASSWORD_LENGTH:
            raise AuthError("password must be at least 12 characters")
        if len(password.encode("utf-8")) > 72:
            raise AuthError("password is too long for the configured bcrypt policy")
        return password

    def register(self, username: str, password: str) -> dict[str, str]:
        user = self._validate_user(username)
        password = self._validate_password(password)
        if user in self._passwords:
            raise AuthError("username already exists")
        self._passwords[user] = bcrypt.hashpw(password.encode(), bcrypt.gensalt(rounds=12))
        return {"user": user}

    def login(self, username: str, password: str, *, issue_jwt: bool = False) -> dict[str, str]:
        user = self._validate_user(username)
        if not isinstance(password, str):
            raise AuthError("invalid credentials")
        stored = self._passwords.get(user)
        if stored is None or not bcrypt.checkpw(password.encode(), stored):
            raise AuthError("invalid credentials")
        token = secrets.token_urlsafe(32)
        csrf = secrets.token_urlsafe(24)
        self._sessions[token] = Session(user=user, csrf=csrf, expires_at=time.time() + self._session_ttl)
        result = {"session_token": token, "csrf_token": csrf, "user": user}
        if issue_jwt:
            result["jwt"] = self.issue_jwt(user)
        return result

    def issue_jwt(self, username: str, *, ttl: int | None = None) -> str:
        user = self._validate_user(username)
        now = int(time.time())
        header = _b64(json.dumps({"alg": "HS256", "typ": "JWT"}, separators=(",", ":")).encode())
        payload = _b64(json.dumps({"sub": user, "iat": now, "exp": now + (ttl or self._session_ttl)}, separators=(",", ":")).encode())
        unsigned = f"{header}.{payload}".encode()
        return f"{header}.{payload}.{_b64(hmac.new(self._jwt_key, unsigned, hashlib.sha256).digest())}"

    def verify_jwt(self, token: str) -> str:
        try:
            header, payload, signature = token.split(".")
            expected = _b64(hmac.new(self._jwt_key, f"{header}.{payload}".encode(), hashlib.sha256).digest())
            if not hmac.compare_digest(signature, expected):
                raise AuthError("invalid token")
            data: Any = json.loads(_unb64(payload))
            if not isinstance(data, dict) or int(data["exp"]) <= int(time.time()):
                raise AuthError("token expired")
            return self._validate_user(data["sub"])
        except (AuthError, KeyError, TypeError, ValueError, json.JSONDecodeError):
            raise AuthError("invalid token")

    def authenticate(self, session_token: str) -> Session:
        session = self._sessions.get(session_token)
        if session is None or session.expires_at <= time.time():
            self._sessions.pop(session_token, None)
            raise AuthError("session expired or invalid")
        return session

    def require_csrf(self, session_token: str, csrf_token: str) -> Session:
        session = self.authenticate(session_token)
        if not isinstance(csrf_token, str) or not hmac.compare_digest(session.csrf, csrf_token):
            raise AuthError("CSRF validation failed")
        return session

    def logout(self, session_token: str, csrf_token: str) -> None:
        self.require_csrf(session_token, csrf_token)
        self._sessions.pop(session_token, None)
