74 lines
2.0 KiB
Python
74 lines
2.0 KiB
Python
import hashlib
|
|
import logging
|
|
import secrets
|
|
from typing import Any, Literal
|
|
|
|
import jwt
|
|
from argon2 import PasswordHasher
|
|
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
|
|
from zxcvbn import zxcvbn
|
|
|
|
from config import cfg
|
|
from schemas.dto import KeyPair
|
|
from schemas.jwt import UserJWTPayload
|
|
from schemas.providers import ProvidersType
|
|
|
|
ctx = PasswordHasher()
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def hash_password(plain_password: str) -> ...:
|
|
return ctx.hash(plain_password)
|
|
|
|
|
|
def verify_password(hashed_password: str, plain_password: str) -> bool:
|
|
try:
|
|
ctx.verify(hashed_password, plain_password)
|
|
return True
|
|
except (VerifyMismatchError, VerificationError, InvalidHashError):
|
|
return False
|
|
except Exception:
|
|
logger.exception("unexpected error while comparing password and hash")
|
|
return False
|
|
|
|
|
|
def generate_jwt(payload: dict[str, Any]) -> str:
|
|
return jwt.encode(payload, cfg.private_key, "EdDSA")
|
|
|
|
|
|
def get_kid_from_token(token: str) -> str | None:
|
|
return jwt.get_unverified_header(token).get("kid")
|
|
|
|
|
|
def decode_jwt(
|
|
token: str, public_key: str, algo: Literal["EdDSA"] = "EdDSA"
|
|
) -> dict[str, Any] | None:
|
|
try:
|
|
return jwt.decode(token, public_key, algo)
|
|
except jwt.ExpiredSignatureError:
|
|
return
|
|
|
|
|
|
def decode_user_jwt(token: str) -> dict[str, Any] | None:
|
|
return decode_jwt(token, cfg.public_key, algo="EdDSA")
|
|
|
|
|
|
def generate_pair(user_id: int, iss: ProvidersType) -> KeyPair:
|
|
payload = UserJWTPayload(sub=str(user_id), iss=iss)
|
|
access_token = generate_jwt(payload.model_dump())
|
|
refresh_token = secrets.token_urlsafe(32)
|
|
|
|
return KeyPair(access_token=access_token, refresh_token=refresh_token, expires_at=payload.exp)
|
|
|
|
|
|
def hash_refresh_token(token: str):
|
|
return hashlib.sha256(token.encode()).hexdigest()
|
|
|
|
|
|
def estimate_password_strength(password: str) -> bool:
|
|
if len(password) < cfg.min_password_length:
|
|
return False
|
|
|
|
r = zxcvbn(password)
|
|
return r.get("score", 0) > cfg.password_security_threshold
|