"""
資安核心:密碼雜湊 + JWT 簽發/驗證。

設計重點:
  1. 密碼用 bcrypt(慢、抗暴力破解、業界標準)
  2. JWT 含 sub(使用者 id)、type(access / refresh)、exp(到期時間)
  3. access / refresh 用同一把密鑰簽,但 type 欄位區分,避免互相混用
"""

from datetime import datetime, timedelta, timezone
from typing import Literal
from uuid import UUID

import jwt
from passlib.context import CryptContext

from app.config import settings


# ---------------------------------------------------------------------------
# 密碼雜湊
# ---------------------------------------------------------------------------
# bcrypt 是慢雜湊演算法,即使資料庫被偷,要從 hash 還原密碼也非常耗時
_pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")


def hash_password(plain_password: str) -> str:
    """把明文密碼轉成 bcrypt 雜湊字串(資料庫存這個)。"""
    return _pwd_context.hash(plain_password)


def verify_password(plain_password: str, password_hash: str) -> bool:
    """驗證明文密碼是否符合資料庫的雜湊。"""
    return _pwd_context.verify(plain_password, password_hash)


# ---------------------------------------------------------------------------
# JWT
# ---------------------------------------------------------------------------
TokenType = Literal["access", "refresh"]


def _create_token(user_id: UUID, token_type: TokenType, expires_delta: timedelta) -> str:
    """簽發 JWT 的內部共用函式。"""
    now = datetime.now(timezone.utc)
    payload = {
        "sub": str(user_id),    # subject = 使用者 ID
        "type": token_type,     # access / refresh
        "iat": int(now.timestamp()),                            # issued at
        "exp": int((now + expires_delta).timestamp()),          # expires at
    }
    return jwt.encode(payload, settings.JWT_SECRET_KEY, algorithm=settings.JWT_ALGORITHM)


def create_access_token(user_id: UUID) -> str:
    """簽一個 access token(短效,用來呼叫 API)。"""
    return _create_token(
        user_id,
        token_type="access",
        expires_delta=timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES),
    )


def create_refresh_token(user_id: UUID) -> str:
    """簽一個 refresh token(長效,用來換新的 access token)。"""
    return _create_token(
        user_id,
        token_type="refresh",
        expires_delta=timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS),
    )


class TokenError(Exception):
    """JWT 驗證失敗(過期、簽章錯、type 不對)時丟出。"""


def decode_token(token: str, expected_type: TokenType) -> UUID:
    """驗證並解析 JWT,回傳 user_id。

    Args:
        token         : JWT 字串
        expected_type : 期望的 token 類型("access" 或 "refresh")
                        例如保護 API 的 endpoint 只能收 access token,
                        /auth/refresh 只能收 refresh token,避免亂用。

    Raises:
        TokenError: token 過期、簽章不對、type 不符等
    """
    try:
        payload = jwt.decode(token, settings.JWT_SECRET_KEY, algorithms=[settings.JWT_ALGORITHM])
    except jwt.ExpiredSignatureError as e:
        raise TokenError("token 已過期") from e
    except jwt.InvalidTokenError as e:
        raise TokenError("無效的 token") from e

    if payload.get("type") != expected_type:
        raise TokenError(f"token 類型錯誤,需要 {expected_type}")

    sub = payload.get("sub")
    if not sub:
        raise TokenError("token 缺少 sub 欄位")

    try:
        return UUID(sub)
    except ValueError as e:
        raise TokenError("token sub 不是合法 UUID") from e
