feat: add register/login/logout with revocable server-side sessions

Claude-Session: https://claude.ai/code/session_01M9MLit5Ko3X4s7rzC5Kv7X
This commit is contained in:
2026-10-02 14:15:24 -05:00
parent 7c09154b76
commit d8d71bf942
2 changed files with 234 additions and 1 deletions
+127 -1
View File
@@ -5,10 +5,15 @@ import hashlib
import hmac import hmac
import logging import logging
import os import os
import re
import secrets import secrets
from pathlib import Path 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__) log = logging.getLogger(__name__)
bp = Blueprint("auth", __name__, url_prefix="/api/auth") bp = Blueprint("auth", __name__, url_prefix="/api/auth")
@@ -16,6 +21,10 @@ bp = Blueprint("auth", __name__, url_prefix="/api/auth")
ITERATIONS = 600_000 ITERATIONS = 600_000
SALT_BYTES = 16 SALT_BYTES = 16
MIN_PEPPER_CHARS = 32 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: 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)." "the file to regenerate it (existing passwords will stop verifying)."
) )
return pepper.encode() 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,)))
+107
View File
@@ -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$"))