diff --git a/backend/auth.py b/backend/auth.py index b70839a..5f9a34c 100644 --- a/backend/auth.py +++ b/backend/auth.py @@ -5,10 +5,15 @@ import hashlib import hmac import logging import os +import re import secrets from pathlib import Path -from flask import Blueprint, Flask +import psycopg2.errors +from flask import Blueprint, Flask, current_app, g, jsonify, make_response, request + +from db import query, query_one +from validation import ApiError, json_body, require_utf8 log = logging.getLogger(__name__) bp = Blueprint("auth", __name__, url_prefix="/api/auth") @@ -16,6 +21,10 @@ bp = Blueprint("auth", __name__, url_prefix="/api/auth") ITERATIONS = 600_000 SALT_BYTES = 16 MIN_PEPPER_CHARS = 32 +SESSION_COOKIE = "sid" +USERNAME_PATTERN = re.compile(r"[A-Za-z0-9_.-]{3,64}") +MIN_PASSWORD, MAX_PASSWORD = 12, 1024 +INVALID_LOGIN = "Invalid username or password" def init_app(app: Flask) -> None: @@ -85,3 +94,120 @@ def load_pepper(pepper_file: Path) -> bytes: "the file to regenerate it (existing passwords will stop verifying)." ) return pepper.encode() + + +# --- Sessions ------------------------------------------------------------------ + +def login_required(view): + @functools.wraps(view) + def wrapper(*args, **kwargs): + token = request.cookies.get(SESSION_COOKIE) + row = None + if token: + # One statement validates and slides the idle window. The created_at predicate enforces + # the 12-hour cap even while expires_at is still in the future (ADR-0002). + row = query_one( + """ + UPDATE sessions + SET expires_at = LEAST(now() + interval '30 minutes', created_at + interval '12 hours') + WHERE token_hash = %s + AND expires_at > now() + AND created_at > now() - interval '12 hours' + RETURNING user_id + """, + (_token_hash(token),), + ) + if row is None: + raise ApiError(401, "Not signed in or session expired; log in again") + g.user_id = row["user_id"] + return view(*args, **kwargs) + + return wrapper + + +def _token_hash(token: str) -> bytes: + return hashlib.sha256(token.encode("utf-8")).digest() + + +def _pepper() -> bytes: + return current_app.extensions["password_pepper"] + + +def _start_session(user: dict, status: int): + token = secrets.token_urlsafe(32) + query("DELETE FROM sessions WHERE user_id = %s AND expires_at <= now()", (user["id"],)) + query( + "INSERT INTO sessions (token_hash, user_id, expires_at) VALUES (%s, %s, now() + interval '30 minutes')", + (_token_hash(token), user["id"]), + ) + response = jsonify(id=user["id"], username=user["username"]) + response.status_code = status + response.set_cookie(SESSION_COOKIE, token, httponly=True, secure=True, samesite="Lax", path="/api") + return response + + +# --- Routes ---------------------------------------------------------------------- + +def _password(body: dict) -> str: + password = body.get("password") + if not isinstance(password, str): + raise ApiError(400, "password is required and must be a string", "password") + if len(password) > MAX_PASSWORD: + raise ApiError(400, f"password must be at most {MAX_PASSWORD} characters", "password") + require_utf8(password, "password") + return password + + +@bp.post("/register") +def register(): + body = json_body({"username", "password"}) + username = body.get("username") + if not isinstance(username, str) or not USERNAME_PATTERN.fullmatch(username): + raise ApiError(400, "username must be 3-64 characters: letters, digits, '.', '_' or '-'", "username") + password = _password(body) + if len(password) < MIN_PASSWORD: + raise ApiError(400, f"password must be at least {MIN_PASSWORD} characters", "password") + try: + user = query_one( + "INSERT INTO users (username, password_hash) VALUES (%s, %s) RETURNING id, username", + (username, hash_password(password, _pepper())), + ) + except psycopg2.errors.UniqueViolation: + raise ApiError(409, "Username already taken; choose another", "username") from None + return _start_session(user, 201) + + +@bp.post("/login") +def login(): + body = json_body({"username", "password"}) + password = _password(body) + username = body.get("username") + user = None + if isinstance(username, str) and USERNAME_PATTERN.fullmatch(username): + user = query_one( + "SELECT id, username, password_hash FROM users WHERE lower(username) = lower(%s)", (username,) + ) + if user is None: + verify_password(password, _dummy_hash(), _pepper()) + raise ApiError(401, INVALID_LOGIN) + if not verify_password(password, user["password_hash"], _pepper()): + raise ApiError(401, INVALID_LOGIN) + if needs_rehash(user["password_hash"]): + query("UPDATE users SET password_hash = %s WHERE id = %s", (hash_password(password, _pepper()), user["id"])) + return _start_session(user, 200) + + +@bp.post("/logout") +def logout(): + token = request.cookies.get(SESSION_COOKIE) + if token: + query("DELETE FROM sessions WHERE token_hash = %s", (_token_hash(token),)) + response = make_response("", 204) + response.delete_cookie(SESSION_COOKIE, path="/api", secure=True, httponly=True, samesite="Lax") + return response + + +@bp.get("/me") +@login_required +def me(): + return jsonify(query_one("SELECT id, username FROM users WHERE id = %s", (g.user_id,))) diff --git a/backend/tests/test_auth.py b/backend/tests/test_auth.py new file mode 100644 index 0000000..eb89645 --- /dev/null +++ b/backend/tests/test_auth.py @@ -0,0 +1,107 @@ +import hashlib +import os + +from app import app +from auth import hash_password +from tests.support import ApiTestCase, db_execute + + +class AuthTests(ApiTestCase): + def session_token(self, client=None) -> str: + return (client or self.client).get_cookie("sid", domain="localhost", path="/api").value + + def test_register_logs_in_and_me_returns_user(self): + response = self.register("Alice") + self.assertEqual(response.get_json(), {"id": 1, "username": "Alice"}) + me = self.call("GET", "/api/auth/me") + self.assertEqual(me.status_code, 200) + self.assertEqual(me.get_json(), {"id": 1, "username": "Alice"}) + + def test_session_cookie_flags(self): + cookie = self.register().headers["Set-Cookie"] + for flag in ("HttpOnly", "Secure", "SameSite=Lax", "Path=/api"): + self.assertIn(flag, cookie) + + def test_session_token_is_stored_hashed(self): + self.register() + digest = hashlib.sha256(self.session_token().encode()).digest() + self.assertEqual(db_execute("SELECT count(*) FROM sessions WHERE token_hash = %s", (digest,))[0][0], 1) + + def test_duplicate_username_is_case_insensitive(self): + self.register("Alice") + response = self.call("POST", "/api/auth/register", {"username": "alice", "password": "another long password"}) + self.assertEqual(response.status_code, 409) + self.assertEqual(response.get_json()["field"], "username") + + def test_register_validation(self): + cases = [ + ({"username": "al", "password": "long enough password"}, "username"), + ({"username": "al ice", "password": "long enough password"}, "username"), + ({"username": "alice", "password": "short"}, "password"), + ({"username": "alice", "password": "x" * 1025}, "password"), + ({"username": "alice", "password": "\ud800" * 12}, "password"), + ({"username": "alice", "password": "long enough password", "is_admin": True}, "is_admin"), + ] + for payload, field in cases: + with self.subTest(payload=payload): + response = self.call("POST", "/api/auth/register", payload) + self.assertEqual(response.status_code, 400) + self.assertEqual(response.get_json()["field"], field) + + def test_login_is_case_insensitive_on_username(self): + self.register("Alice") + self.client = app.test_client() # fresh cookie jar: signed out + response = self.call("POST", "/api/auth/login", {"username": "ALICE", "password": "correct horse battery"}) + self.assertEqual(response.status_code, 200) + self.assertEqual(response.get_json(), {"id": 1, "username": "Alice"}) + + def test_wrong_password_and_unknown_user_are_indistinguishable(self): + self.register() + wrong = self.call("POST", "/api/auth/login", {"username": "alice", "password": "not the password"}) + unknown = self.call("POST", "/api/auth/login", {"username": "nobody", "password": "not the password"}) + self.assertEqual(wrong.status_code, 401) + self.assertEqual(unknown.status_code, 401) + self.assertEqual(wrong.get_json(), unknown.get_json()) + self.assertEqual(wrong.get_json(), {"error": "Invalid username or password"}) + + def test_login_with_malformed_username_is_401_not_500(self): + response = self.call("POST", "/api/auth/login", {"username": "a\x00b", "password": "whatever password"}) + self.assertEqual(response.status_code, 401) + + def test_me_without_session_is_401(self): + response = self.call("GET", "/api/auth/me") + self.assertEqual(response.status_code, 401) + self.assertEqual(response.get_json(), {"error": "Not signed in or session expired; log in again"}) + + def test_logout_revokes_token_server_side(self): + self.register() + token = self.session_token() + self.assertEqual(self.call("POST", "/api/auth/logout").status_code, 204) + self.client.set_cookie("sid", token, domain="localhost", path="/api") # replay the stolen cookie + self.assertEqual(self.call("GET", "/api/auth/me").status_code, 401) + + def test_idle_session_expires(self): + self.register() + db_execute("UPDATE sessions SET expires_at = now() - interval '1 second'") + self.assertEqual(self.call("GET", "/api/auth/me").status_code, 401) + + def test_absolute_cap_applies_even_inside_idle_window(self): + self.register() + db_execute("UPDATE sessions SET created_at = now() - interval '12 hours 1 second'") + self.assertEqual(self.call("GET", "/api/auth/me").status_code, 401) + + def test_activity_slides_idle_expiry(self): + self.register() + db_execute("UPDATE sessions SET expires_at = now() + interval '1 minute'") + self.assertEqual(self.call("GET", "/api/auth/me").status_code, 200) + remaining = db_execute("SELECT expires_at - now() > interval '29 minutes' FROM sessions")[0][0] + self.assertTrue(remaining) + + def test_old_hash_is_upgraded_on_login(self): + self.register() + pepper = os.environ["PASSWORD_PEPPER"].encode() + old = hash_password("correct horse battery", pepper, iterations=1000) + db_execute("UPDATE users SET password_hash = %s", (old,)) + response = self.call("POST", "/api/auth/login", {"username": "alice", "password": "correct horse battery"}) + self.assertEqual(response.status_code, 200) + self.assertTrue(db_execute("SELECT password_hash FROM users")[0][0].startswith("pbkdf2_sha256$600000$"))