88 lines
2.5 KiB
Python
88 lines
2.5 KiB
Python
from __future__ import annotations
|
|
|
|
import base64
|
|
import hashlib
|
|
import hmac
|
|
import json
|
|
import os
|
|
import time
|
|
from dataclasses import dataclass
|
|
|
|
from app.core.config import get_settings
|
|
|
|
settings = get_settings()
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class SessionPayload:
|
|
user_id: str
|
|
public_ref: str
|
|
role: str
|
|
display_name: str
|
|
issued_at: int
|
|
session_id: str = ""
|
|
|
|
|
|
def _sign(data: bytes) -> str:
|
|
digest = hmac.new(settings.app_secret.encode(), data, hashlib.sha256).digest()
|
|
return base64.urlsafe_b64encode(digest).decode().rstrip("=")
|
|
|
|
|
|
def session_token_hash(token: str) -> str:
|
|
"""Return a non-reversible identifier safe to persist for token revocation."""
|
|
return hashlib.sha256(token.encode()).hexdigest()
|
|
|
|
|
|
def create_session_token(payload: SessionPayload) -> str:
|
|
body = json.dumps(payload.__dict__, separators=(",", ":")).encode()
|
|
encoded_body = base64.urlsafe_b64encode(body).decode().rstrip("=")
|
|
signature = _sign(encoded_body.encode())
|
|
return f"{encoded_body}.{signature}"
|
|
|
|
|
|
def read_session_token(token: str) -> SessionPayload | None:
|
|
try:
|
|
encoded_body, signature = token.split(".", 1)
|
|
except ValueError:
|
|
return None
|
|
expected = _sign(encoded_body.encode())
|
|
if not hmac.compare_digest(expected, signature):
|
|
return None
|
|
padding = "=" * (-len(encoded_body) % 4)
|
|
try:
|
|
body = json.loads(base64.urlsafe_b64decode(encoded_body + padding))
|
|
except (ValueError, json.JSONDecodeError):
|
|
return None
|
|
payload = SessionPayload(**body)
|
|
if time.time() - payload.issued_at > settings.session_ttl_seconds:
|
|
return None
|
|
return payload
|
|
|
|
|
|
def hash_password(password: str) -> str:
|
|
salt = os.urandom(16)
|
|
derived = hashlib.scrypt(password.encode(), salt=salt, n=2**14, r=8, p=1, dklen=32)
|
|
encoded_salt = base64.urlsafe_b64encode(salt).decode()
|
|
encoded_hash = base64.urlsafe_b64encode(derived).decode()
|
|
return f"scrypt$16384$8$1${encoded_salt}${encoded_hash}"
|
|
|
|
|
|
def verify_password(password: str, encoded: str | None) -> bool:
|
|
if not encoded:
|
|
return False
|
|
try:
|
|
algorithm, n, r, p, salt, expected = encoded.split("$")
|
|
if algorithm != "scrypt":
|
|
return False
|
|
derived = hashlib.scrypt(
|
|
password.encode(),
|
|
salt=base64.urlsafe_b64decode(salt.encode()),
|
|
n=int(n),
|
|
r=int(r),
|
|
p=int(p),
|
|
dklen=32,
|
|
)
|
|
return hmac.compare_digest(derived, base64.urlsafe_b64decode(expected.encode()))
|
|
except (ValueError, TypeError):
|
|
return False
|