style: format backend with black

Also drops the session-sliding comment in login_required.
This commit is contained in:
2026-10-02 14:46:16 -05:00
parent 2512c67893
commit 63c9991421
14 changed files with 320 additions and 82 deletions
+20 -4
View File
@@ -1,4 +1,5 @@
"""Application entry point: wires slices, cross-cutting request rules, and JSON error handling.""" """Application entry point: wires slices, cross-cutting request rules, and JSON error handling."""
import logging import logging
import psycopg2.errors import psycopg2.errors
@@ -16,7 +17,9 @@ MUTATING_METHODS = {"POST", "PUT", "PATCH", "DELETE"}
def create_app() -> Flask: def create_app() -> Flask:
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s") logging.basicConfig(
level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s"
)
app = Flask(__name__) app = Flask(__name__)
# Werkzeug buffers and parses the whole body before our validation runs, so unauthenticated # Werkzeug buffers and parses the whole body before our validation runs, so unauthenticated
# callers could exhaust memory; 1 MiB is far above the largest legitimate payload (10,000-char note). # callers could exhaust memory; 1 MiB is far above the largest legitimate payload (10,000-char note).
@@ -35,11 +38,19 @@ def create_app() -> Flask:
# CSRF defense (ADR-0002): cross-site requests can only send this content type after a # CSRF defense (ADR-0002): cross-site requests can only send this content type after a
# CORS preflight, and this API never grants CORS. # CORS preflight, and this API never grants CORS.
if request.method in MUTATING_METHODS and not request.is_json: if request.method in MUTATING_METHODS and not request.is_json:
raise ApiError(400, "Request body must be JSON with Content-Type: application/json") raise ApiError(
400, "Request body must be JSON with Content-Type: application/json"
)
@app.after_request @app.after_request
def log_request(response): def log_request(response):
log.info("%s %s -> %s user=%s", request.method, request.path, response.status_code, g.get("user_id")) log.info(
"%s %s -> %s user=%s",
request.method,
request.path,
response.status_code,
g.get("user_id"),
)
return response return response
@app.errorhandler(ApiError) @app.errorhandler(ApiError)
@@ -55,7 +66,12 @@ def create_app() -> Flask:
@app.errorhandler(psycopg2.errors.CheckViolation) @app.errorhandler(psycopg2.errors.CheckViolation)
def handle_check_violation(error: psycopg2.errors.CheckViolation): def handle_check_violation(error: psycopg2.errors.CheckViolation):
return jsonify(error=f"Value breaks data rule '{error.diag.constraint_name}'; correct it and retry"), 400 return (
jsonify(
error=f"Value breaks data rule '{error.diag.constraint_name}'; correct it and retry"
),
400,
)
@app.errorhandler(Exception) @app.errorhandler(Exception)
def handle_unexpected(_error: Exception): def handle_unexpected(_error: Exception):
+43 -13
View File
@@ -1,4 +1,5 @@
"""Auth slice: password hashing (ADR-0001), server-side sessions (ADR-0002), and auth routes.""" """Auth slice: password hashing (ADR-0001), server-side sessions (ADR-0002), and auth routes."""
import base64 import base64
import functools import functools
import hashlib import hashlib
@@ -35,7 +36,14 @@ def init_app(app: Flask) -> None:
# --- Passwords ----------------------------------------------------------------- # --- Passwords -----------------------------------------------------------------
def hash_password(password: str, pepper: bytes, *, salt: bytes | None = None, iterations: int = ITERATIONS) -> str:
def hash_password(
password: str,
pepper: bytes,
*,
salt: bytes | None = None,
iterations: int = ITERATIONS,
) -> str:
salt = secrets.token_bytes(SALT_BYTES) if salt is None else salt salt = secrets.token_bytes(SALT_BYTES) if salt is None else salt
derived = _derive(password, pepper, salt, iterations) derived = _derive(password, pepper, salt, iterations)
return f"pbkdf2_sha256${iterations}${_b64(salt)}${_b64(derived)}" return f"pbkdf2_sha256${iterations}${_b64(salt)}${_b64(derived)}"
@@ -98,14 +106,13 @@ def load_pepper(pepper_file: Path) -> bytes:
# --- Sessions ------------------------------------------------------------------ # --- Sessions ------------------------------------------------------------------
def login_required(view): def login_required(view):
@functools.wraps(view) @functools.wraps(view)
def wrapper(*args, **kwargs): def wrapper(*args, **kwargs):
token = request.cookies.get(SESSION_COOKIE) token = request.cookies.get(SESSION_COOKIE)
row = None row = None
if token: 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( row = query_one(
""" """
UPDATE sessions UPDATE sessions
@@ -135,25 +142,32 @@ def _pepper() -> bytes:
def _start_session(user: dict, status: int): def _start_session(user: dict, status: int):
token = secrets.token_urlsafe(32) token = secrets.token_urlsafe(32)
query("DELETE FROM sessions WHERE user_id = %s AND expires_at <= now()", (user["id"],)) query(
"DELETE FROM sessions WHERE user_id = %s AND expires_at <= now()", (user["id"],)
)
query( query(
"INSERT INTO sessions (token_hash, user_id, expires_at) VALUES (%s, %s, now() + interval '30 minutes')", "INSERT INTO sessions (token_hash, user_id, expires_at) VALUES (%s, %s, now() + interval '30 minutes')",
(_token_hash(token), user["id"]), (_token_hash(token), user["id"]),
) )
response = jsonify(id=user["id"], username=user["username"]) response = jsonify(id=user["id"], username=user["username"])
response.status_code = status response.status_code = status
response.set_cookie(SESSION_COOKIE, token, httponly=True, secure=True, samesite="Lax", path="/api") response.set_cookie(
SESSION_COOKIE, token, httponly=True, secure=True, samesite="Lax", path="/api"
)
return response return response
# --- Routes ---------------------------------------------------------------------- # --- Routes ----------------------------------------------------------------------
def _password(body: dict) -> str: def _password(body: dict) -> str:
password = body.get("password") password = body.get("password")
if not isinstance(password, str): if not isinstance(password, str):
raise ApiError(400, "password is required and must be a string", "password") raise ApiError(400, "password is required and must be a string", "password")
if len(password) > MAX_PASSWORD: if len(password) > MAX_PASSWORD:
raise ApiError(400, f"password must be at most {MAX_PASSWORD} characters", "password") raise ApiError(
400, f"password must be at most {MAX_PASSWORD} characters", "password"
)
require_utf8(password, "password") require_utf8(password, "password")
return password return password
@@ -163,17 +177,25 @@ def register():
body = json_body({"username", "password"}) body = json_body({"username", "password"})
username = body.get("username") username = body.get("username")
if not isinstance(username, str) or not USERNAME_PATTERN.fullmatch(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") raise ApiError(
400,
"username must be 3-64 characters: letters, digits, '.', '_' or '-'",
"username",
)
password = _password(body) password = _password(body)
if len(password) < MIN_PASSWORD: if len(password) < MIN_PASSWORD:
raise ApiError(400, f"password must be at least {MIN_PASSWORD} characters", "password") raise ApiError(
400, f"password must be at least {MIN_PASSWORD} characters", "password"
)
try: try:
user = query_one( user = query_one(
"INSERT INTO users (username, password_hash) VALUES (%s, %s) RETURNING id, username", "INSERT INTO users (username, password_hash) VALUES (%s, %s) RETURNING id, username",
(username, hash_password(password, _pepper())), (username, hash_password(password, _pepper())),
) )
except psycopg2.errors.UniqueViolation: except psycopg2.errors.UniqueViolation:
raise ApiError(409, "Username already taken; choose another", "username") from None raise ApiError(
409, "Username already taken; choose another", "username"
) from None
return _start_session(user, 201) return _start_session(user, 201)
@@ -185,7 +207,8 @@ def login():
user = None user = None
if isinstance(username, str) and USERNAME_PATTERN.fullmatch(username): if isinstance(username, str) and USERNAME_PATTERN.fullmatch(username):
user = query_one( user = query_one(
"SELECT id, username, password_hash FROM users WHERE lower(username) = lower(%s)", (username,) "SELECT id, username, password_hash FROM users WHERE lower(username) = lower(%s)",
(username,),
) )
if user is None: if user is None:
verify_password(password, _dummy_hash(), _pepper()) verify_password(password, _dummy_hash(), _pepper())
@@ -193,7 +216,10 @@ def login():
if not verify_password(password, user["password_hash"], _pepper()): if not verify_password(password, user["password_hash"], _pepper()):
raise ApiError(401, INVALID_LOGIN) raise ApiError(401, INVALID_LOGIN)
if needs_rehash(user["password_hash"]): if needs_rehash(user["password_hash"]):
query("UPDATE users SET password_hash = %s WHERE id = %s", (hash_password(password, _pepper()), user["id"])) query(
"UPDATE users SET password_hash = %s WHERE id = %s",
(hash_password(password, _pepper()), user["id"]),
)
return _start_session(user, 200) return _start_session(user, 200)
@@ -203,11 +229,15 @@ def logout():
if token: if token:
query("DELETE FROM sessions WHERE token_hash = %s", (_token_hash(token),)) query("DELETE FROM sessions WHERE token_hash = %s", (_token_hash(token),))
response = make_response("", 204) response = make_response("", 204)
response.delete_cookie(SESSION_COOKIE, path="/api", secure=True, httponly=True, samesite="Lax") response.delete_cookie(
SESSION_COOKIE, path="/api", secure=True, httponly=True, samesite="Lax"
)
return response return response
@bp.get("/me") @bp.get("/me")
@login_required @login_required
def me(): def me():
return jsonify(query_one("SELECT id, username FROM users WHERE id = %s", (g.user_id,))) return jsonify(
query_one("SELECT id, username FROM users WHERE id = %s", (g.user_id,))
)
+43 -10
View File
@@ -1,4 +1,5 @@
"""Books slice: CRUD, progress, search/filter, and the genre list.""" """Books slice: CRUD, progress, search/filter, and the genre list."""
import re import re
import psycopg2.errors import psycopg2.errors
@@ -60,7 +61,9 @@ def escape_like(term: str) -> str:
def owned_book(book_id: int) -> dict: def owned_book(book_id: int) -> dict:
row = query_one(BOOK_SELECT + " WHERE b.id = %s AND b.user_id = %s", (book_id, g.user_id)) row = query_one(
BOOK_SELECT + " WHERE b.id = %s AND b.user_id = %s", (book_id, g.user_id)
)
if row is None: if row is None:
raise ApiError(404, "Book not found") raise ApiError(404, "Book not found")
return row return row
@@ -80,7 +83,9 @@ def _save(sql: str, params: tuple) -> dict | None:
try: try:
return query_one(sql, params) return query_one(sql, params)
except psycopg2.errors.ForeignKeyViolation: except psycopg2.errors.ForeignKeyViolation:
raise ApiError(400, "Unknown genre_id; see GET /api/genres", "genre_id") from None raise ApiError(
400, "Unknown genre_id; see GET /api/genres", "genre_id"
) from None
@bp.get("/genres") @bp.get("/genres")
@@ -97,14 +102,18 @@ def list_books():
search = request.args.get("q", "").strip() search = request.args.get("q", "").strip()
if search: if search:
if len(search) > MAX_TEXT or "\x00" in search: if len(search) > MAX_TEXT or "\x00" in search:
raise ApiError(400, f"q must be plain text of at most {MAX_TEXT} characters", "q") raise ApiError(
400, f"q must be plain text of at most {MAX_TEXT} characters", "q"
)
pattern = f"%{escape_like(search)}%" pattern = f"%{escape_like(search)}%"
sql += " AND (b.title ILIKE %s OR b.author ILIKE %s)" sql += " AND (b.title ILIKE %s OR b.author ILIKE %s)"
params += [pattern, pattern] params += [pattern, pattern]
genre_id = request.args.get("genre_id", "") genre_id = request.args.get("genre_id", "")
if genre_id: if genre_id:
if not GENRE_ID_PATTERN.fullmatch(genre_id) or int(genre_id) > MAX_GENRE_ID: if not GENRE_ID_PATTERN.fullmatch(genre_id) or int(genre_id) > MAX_GENRE_ID:
raise ApiError(400, "genre_id must be an id from GET /api/genres", "genre_id") raise ApiError(
400, "genre_id must be an id from GET /api/genres", "genre_id"
)
sql += " AND b.genre_id = %s" sql += " AND b.genre_id = %s"
params.append(int(genre_id)) params.append(int(genre_id))
sql += " ORDER BY b.updated_at DESC, b.id DESC" sql += " ORDER BY b.updated_at DESC, b.id DESC"
@@ -116,7 +125,9 @@ def list_books():
def create_book(): def create_book():
body = json_body(FIELDS) body = json_body(FIELDS)
book = {field: VALIDATORS[field](body) for field in REQUIRED_FIELDS} book = {field: VALIDATORS[field](body) for field in REQUIRED_FIELDS}
book["current_page"] = VALIDATORS["current_page"](body) if "current_page" in body else 0 book["current_page"] = (
VALIDATORS["current_page"](body) if "current_page" in body else 0
)
_check_progress(book, "current_page") _check_progress(book, "current_page")
row = _save( row = _save(
""" """
@@ -124,7 +135,14 @@ def create_book():
VALUES (%s, %s, %s, %s, %s, %s) VALUES (%s, %s, %s, %s, %s, %s)
RETURNING id RETURNING id
""", """,
(g.user_id, book["title"], book["author"], book["genre_id"], book["total_pages"], book["current_page"]), (
g.user_id,
book["title"],
book["author"],
book["genre_id"],
book["total_pages"],
book["current_page"],
),
) )
return jsonify(to_json(owned_book(row["id"]))), 201 return jsonify(to_json(owned_book(row["id"]))), 201
@@ -143,7 +161,9 @@ def update_book(book_id: int):
if not body: if not body:
raise ApiError(400, f"Provide at least one of: {', '.join(sorted(FIELDS))}") raise ApiError(400, f"Provide at least one of: {', '.join(sorted(FIELDS))}")
current = owned_book(book_id) current = owned_book(book_id)
book = {field: current[field] for field in FIELDS} | {field: VALIDATORS[field](body) for field in body} book = {field: current[field] for field in FIELDS} | {
field: VALIDATORS[field](body) for field in body
}
_check_progress(book, "total_pages" if "total_pages" in body else "current_page") _check_progress(book, "total_pages" if "total_pages" in body else "current_page")
_save( _save(
""" """
@@ -151,8 +171,15 @@ def update_book(book_id: int):
SET title = %s, author = %s, genre_id = %s, total_pages = %s, current_page = %s, updated_at = now() SET title = %s, author = %s, genre_id = %s, total_pages = %s, current_page = %s, updated_at = now()
WHERE id = %s AND user_id = %s WHERE id = %s AND user_id = %s
""", """,
(book["title"], book["author"], book["genre_id"], book["total_pages"], book["current_page"], (
book_id, g.user_id), book["title"],
book["author"],
book["genre_id"],
book["total_pages"],
book["current_page"],
book_id,
g.user_id,
),
) )
return jsonify(to_json(owned_book(book_id))) return jsonify(to_json(owned_book(book_id)))
@@ -160,6 +187,12 @@ def update_book(book_id: int):
@bp.delete("/books/<int(max=2147483647):book_id>") @bp.delete("/books/<int(max=2147483647):book_id>")
@login_required @login_required
def delete_book(book_id: int): def delete_book(book_id: int):
if query_one("DELETE FROM books WHERE id = %s AND user_id = %s RETURNING id", (book_id, g.user_id)) is None: if (
query_one(
"DELETE FROM books WHERE id = %s AND user_id = %s RETURNING id",
(book_id, g.user_id),
)
is None
):
raise ApiError(404, "Book not found") raise ApiError(404, "Book not found")
return "", 204 return "", 204
+4 -1
View File
@@ -1,4 +1,5 @@
"""PostgreSQL access: a connection pool, one connection per request, and two query helpers.""" """PostgreSQL access: a connection pool, one connection per request, and two query helpers."""
import os import os
from pathlib import Path from pathlib import Path
@@ -11,7 +12,9 @@ SCHEMA = Path(__file__).with_name("schema.sql")
def init_app(app: Flask) -> None: def init_app(app: Flask) -> None:
pool = psycopg2.pool.ThreadedConnectionPool(1, 5, os.environ.get("DATABASE_URL", DEFAULT_DATABASE_URL)) pool = psycopg2.pool.ThreadedConnectionPool(
1, 5, os.environ.get("DATABASE_URL", DEFAULT_DATABASE_URL)
)
conn = pool.getconn() conn = pool.getconn()
try: try:
with conn, conn.cursor() as cur: with conn, conn.cursor() as cur:
+11 -2
View File
@@ -1,4 +1,5 @@
"""Notes slice: a per-book reading journal. Ownership is always derived through books.user_id.""" """Notes slice: a per-book reading journal. Ownership is always derived through books.user_id."""
from flask import Blueprint, g, jsonify from flask import Blueprint, g, jsonify
from auth import login_required from auth import login_required
@@ -22,7 +23,12 @@ def to_json(row: dict) -> dict:
def _require_book(book_id: int) -> None: def _require_book(book_id: int) -> None:
if query_one("SELECT 1 FROM books WHERE id = %s AND user_id = %s", (book_id, g.user_id)) is None: if (
query_one(
"SELECT 1 FROM books WHERE id = %s AND user_id = %s", (book_id, g.user_id)
)
is None
):
raise ApiError(404, "Book not found") raise ApiError(404, "Book not found")
@@ -42,7 +48,10 @@ def list_notes(book_id: int):
def add_note(book_id: int): def add_note(book_id: int):
_require_book(book_id) _require_book(book_id)
body = text(json_body({"body"}), "body", MAX_BODY) body = text(json_body({"body"}), "body", MAX_BODY)
row = query_one(f"INSERT INTO notes AS n (book_id, body) VALUES (%s, %s) RETURNING {NOTE_COLUMNS}", (book_id, body)) row = query_one(
f"INSERT INTO notes AS n (book_id, body) VALUES (%s, %s) RETURNING {NOTE_COLUMNS}",
(book_id, body),
)
return jsonify(to_json(row)), 201 return jsonify(to_json(row)), 201
+10 -3
View File
@@ -1,10 +1,13 @@
"""Point the app at a dedicated test database before any test module imports it.""" """Point the app at a dedicated test database before any test module imports it."""
import os import os
import psycopg2 import psycopg2
from psycopg2 import sql from psycopg2 import sql
ADMIN_URL = os.environ.get("TEST_ADMIN_DATABASE_URL", "postgresql://postgres:[email protected]:5432/postgres") ADMIN_URL = os.environ.get(
"TEST_ADMIN_DATABASE_URL", "postgresql://postgres:[email protected]:5432/postgres"
)
TEST_DATABASE = "books_test" TEST_DATABASE = "books_test"
@@ -13,9 +16,13 @@ def _ensure_test_database() -> None:
conn.autocommit = True # CREATE DATABASE cannot run inside a transaction conn.autocommit = True # CREATE DATABASE cannot run inside a transaction
try: try:
with conn.cursor() as cur: with conn.cursor() as cur:
cur.execute("SELECT 1 FROM pg_database WHERE datname = %s", (TEST_DATABASE,)) cur.execute(
"SELECT 1 FROM pg_database WHERE datname = %s", (TEST_DATABASE,)
)
if cur.fetchone() is None: if cur.fetchone() is None:
cur.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(TEST_DATABASE))) cur.execute(
sql.SQL("CREATE DATABASE {}").format(sql.Identifier(TEST_DATABASE))
)
finally: finally:
conn.close() conn.close()
+6 -1
View File
@@ -34,7 +34,12 @@ class ApiTestCase(unittest.TestCase):
return (client or self.client).open(path, **kwargs) return (client or self.client).open(path, **kwargs)
def register(self, username="alice", password="correct horse battery", client=None): def register(self, username="alice", password="correct horse battery", client=None):
response = self.call("POST", "/api/auth/register", {"username": username, "password": password}, client) response = self.call(
"POST",
"/api/auth/register",
{"username": username, "password": password},
client,
)
self.assertEqual(response.status_code, 201, response.get_json()) self.assertEqual(response.status_code, 201, response.get_json())
return response return response
+11 -3
View File
@@ -16,7 +16,9 @@ class AppTests(ApiTestCase):
db_execute(SCHEMA.read_text()) db_execute(SCHEMA.read_text())
db_execute(SCHEMA.read_text()) db_execute(SCHEMA.read_text())
self.assertEqual(db_execute("SELECT count(*) FROM genres")[0][0], 16) self.assertEqual(db_execute("SELECT count(*) FROM genres")[0][0], 16)
self.assertEqual(db_execute("SELECT last_value FROM genres_id_seq")[0][0], before) self.assertEqual(
db_execute("SELECT last_value FROM genres_id_seq")[0][0], before
)
def test_unknown_route_returns_json_404(self): def test_unknown_route_returns_json_404(self):
response = self.call("GET", "/api/nope") response = self.call("GET", "/api/nope")
@@ -24,7 +26,9 @@ class AppTests(ApiTestCase):
self.assertIn("error", response.get_json()) self.assertIn("error", response.get_json())
def test_mutation_without_json_content_type_is_rejected(self): def test_mutation_without_json_content_type_is_rejected(self):
response = self.client.post("/api/health", data="x", content_type="text/plain", base_url=HTTPS) response = self.client.post(
"/api/health", data="x", content_type="text/plain", base_url=HTTPS
)
self.assertEqual(response.status_code, 400) self.assertEqual(response.status_code, 400)
self.assertEqual( self.assertEqual(
response.get_json(), response.get_json(),
@@ -32,6 +36,10 @@ class AppTests(ApiTestCase):
) )
def test_oversized_body_is_rejected_before_parsing(self): def test_oversized_body_is_rejected_before_parsing(self):
response = self.call("POST", "/api/auth/login", {"username": "a", "password": "x" * (2 * 1024 * 1024)}) response = self.call(
"POST",
"/api/auth/login",
{"username": "a", "password": "x" * (2 * 1024 * 1024)},
)
self.assertEqual(response.status_code, 413) self.assertEqual(response.status_code, 413)
self.assertIn("error", response.get_json()) self.assertIn("error", response.get_json())
+67 -14
View File
@@ -8,7 +8,11 @@ from tests.support import ApiTestCase, db_execute
class AuthTests(ApiTestCase): class AuthTests(ApiTestCase):
def session_token(self, client=None) -> str: def session_token(self, client=None) -> str:
return (client or self.client).get_cookie("sid", domain="localhost", path="/api").value return (
(client or self.client)
.get_cookie("sid", domain="localhost", path="/api")
.value
)
def test_register_logs_in_and_me_returns_user(self): def test_register_logs_in_and_me_returns_user(self):
response = self.register("Alice") response = self.register("Alice")
@@ -25,11 +29,20 @@ class AuthTests(ApiTestCase):
def test_session_token_is_stored_hashed(self): def test_session_token_is_stored_hashed(self):
self.register() self.register()
digest = hashlib.sha256(self.session_token().encode()).digest() digest = hashlib.sha256(self.session_token().encode()).digest()
self.assertEqual(db_execute("SELECT count(*) FROM sessions WHERE token_hash = %s", (digest,))[0][0], 1) self.assertEqual(
db_execute(
"SELECT count(*) FROM sessions WHERE token_hash = %s", (digest,)
)[0][0],
1,
)
def test_duplicate_username_is_case_insensitive(self): def test_duplicate_username_is_case_insensitive(self):
self.register("Alice") self.register("Alice")
response = self.call("POST", "/api/auth/register", {"username": "alice", "password": "another long password"}) response = self.call(
"POST",
"/api/auth/register",
{"username": "alice", "password": "another long password"},
)
self.assertEqual(response.status_code, 409) self.assertEqual(response.status_code, 409)
self.assertEqual(response.get_json()["field"], "username") self.assertEqual(response.get_json()["field"], "username")
@@ -40,7 +53,14 @@ class AuthTests(ApiTestCase):
({"username": "alice", "password": "short"}, "password"), ({"username": "alice", "password": "short"}, "password"),
({"username": "alice", "password": "x" * 1025}, "password"), ({"username": "alice", "password": "x" * 1025}, "password"),
({"username": "alice", "password": "\ud800" * 12}, "password"), ({"username": "alice", "password": "\ud800" * 12}, "password"),
({"username": "alice", "password": "long enough password", "is_admin": True}, "is_admin"), (
{
"username": "alice",
"password": "long enough password",
"is_admin": True,
},
"is_admin",
),
] ]
for payload, field in cases: for payload, field in cases:
with self.subTest(payload=payload): with self.subTest(payload=payload):
@@ -51,33 +71,54 @@ class AuthTests(ApiTestCase):
def test_login_is_case_insensitive_on_username(self): def test_login_is_case_insensitive_on_username(self):
self.register("Alice") self.register("Alice")
self.client = app.test_client() # fresh cookie jar: signed out self.client = app.test_client() # fresh cookie jar: signed out
response = self.call("POST", "/api/auth/login", {"username": "ALICE", "password": "correct horse battery"}) response = self.call(
"POST",
"/api/auth/login",
{"username": "ALICE", "password": "correct horse battery"},
)
self.assertEqual(response.status_code, 200) self.assertEqual(response.status_code, 200)
self.assertEqual(response.get_json(), {"id": 1, "username": "Alice"}) self.assertEqual(response.get_json(), {"id": 1, "username": "Alice"})
def test_wrong_password_and_unknown_user_are_indistinguishable(self): def test_wrong_password_and_unknown_user_are_indistinguishable(self):
self.register() self.register()
wrong = self.call("POST", "/api/auth/login", {"username": "alice", "password": "not the password"}) wrong = self.call(
unknown = self.call("POST", "/api/auth/login", {"username": "nobody", "password": "not the password"}) "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(wrong.status_code, 401)
self.assertEqual(unknown.status_code, 401) self.assertEqual(unknown.status_code, 401)
self.assertEqual(wrong.get_json(), unknown.get_json()) self.assertEqual(wrong.get_json(), unknown.get_json())
self.assertEqual(wrong.get_json(), {"error": "Invalid username or password"}) self.assertEqual(wrong.get_json(), {"error": "Invalid username or password"})
def test_login_with_malformed_username_is_401_not_500(self): def test_login_with_malformed_username_is_401_not_500(self):
response = self.call("POST", "/api/auth/login", {"username": "a\x00b", "password": "whatever password"}) response = self.call(
"POST",
"/api/auth/login",
{"username": "a\x00b", "password": "whatever password"},
)
self.assertEqual(response.status_code, 401) self.assertEqual(response.status_code, 401)
def test_me_without_session_is_401(self): def test_me_without_session_is_401(self):
response = self.call("GET", "/api/auth/me") response = self.call("GET", "/api/auth/me")
self.assertEqual(response.status_code, 401) self.assertEqual(response.status_code, 401)
self.assertEqual(response.get_json(), {"error": "Not signed in or session expired; log in again"}) self.assertEqual(
response.get_json(),
{"error": "Not signed in or session expired; log in again"},
)
def test_logout_revokes_token_server_side(self): def test_logout_revokes_token_server_side(self):
self.register() self.register()
token = self.session_token() token = self.session_token()
self.assertEqual(self.call("POST", "/api/auth/logout").status_code, 204) 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.client.set_cookie(
"sid", token, domain="localhost", path="/api"
) # replay the stolen cookie
self.assertEqual(self.call("GET", "/api/auth/me").status_code, 401) self.assertEqual(self.call("GET", "/api/auth/me").status_code, 401)
def test_idle_session_expires(self): def test_idle_session_expires(self):
@@ -87,14 +128,18 @@ class AuthTests(ApiTestCase):
def test_absolute_cap_applies_even_inside_idle_window(self): def test_absolute_cap_applies_even_inside_idle_window(self):
self.register() self.register()
db_execute("UPDATE sessions SET created_at = now() - interval '12 hours 1 second'") db_execute(
"UPDATE sessions SET created_at = now() - interval '12 hours 1 second'"
)
self.assertEqual(self.call("GET", "/api/auth/me").status_code, 401) self.assertEqual(self.call("GET", "/api/auth/me").status_code, 401)
def test_activity_slides_idle_expiry(self): def test_activity_slides_idle_expiry(self):
self.register() self.register()
db_execute("UPDATE sessions SET expires_at = now() + interval '1 minute'") db_execute("UPDATE sessions SET expires_at = now() + interval '1 minute'")
self.assertEqual(self.call("GET", "/api/auth/me").status_code, 200) 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] remaining = db_execute(
"SELECT expires_at - now() > interval '29 minutes' FROM sessions"
)[0][0]
self.assertTrue(remaining) self.assertTrue(remaining)
def test_old_hash_is_upgraded_on_login(self): def test_old_hash_is_upgraded_on_login(self):
@@ -102,6 +147,14 @@ class AuthTests(ApiTestCase):
pepper = os.environ["PASSWORD_PEPPER"].encode() pepper = os.environ["PASSWORD_PEPPER"].encode()
old = hash_password("correct horse battery", pepper, iterations=1000) old = hash_password("correct horse battery", pepper, iterations=1000)
db_execute("UPDATE users SET password_hash = %s", (old,)) db_execute("UPDATE users SET password_hash = %s", (old,))
response = self.call("POST", "/api/auth/login", {"username": "alice", "password": "correct horse battery"}) response = self.call(
"POST",
"/api/auth/login",
{"username": "alice", "password": "correct horse battery"},
)
self.assertEqual(response.status_code, 200) self.assertEqual(response.status_code, 200)
self.assertTrue(db_execute("SELECT password_hash FROM users")[0][0].startswith("pbkdf2_sha256$600000$")) self.assertTrue(
db_execute("SELECT password_hash FROM users")[0][0].startswith(
"pbkdf2_sha256$600000$"
)
)
+48 -13
View File
@@ -9,8 +9,14 @@ class BookTests(ApiTestCase):
def test_create_returns_full_shape(self): def test_create_returns_full_shape(self):
book = self.create_book() book = self.create_book()
self.assertEqual(book["title"], "Dune") self.assertEqual(book["title"], "Dune")
self.assertEqual(book["genre"], {"id": self.genre_id("Science Fiction"), "name": "Science Fiction"}) self.assertEqual(
self.assertEqual((book["current_page"], book["total_pages"], book["status"]), (0, 412, "not_started")) book["genre"],
{"id": self.genre_id("Science Fiction"), "name": "Science Fiction"},
)
self.assertEqual(
(book["current_page"], book["total_pages"], book["status"]),
(0, 412, "not_started"),
)
self.assertIn("T", book["created_at"]) # ISO 8601 self.assertIn("T", book["created_at"]) # ISO 8601
def test_create_trims_text(self): def test_create_trims_text(self):
@@ -19,25 +25,44 @@ class BookTests(ApiTestCase):
def test_progress_updates_status(self): def test_progress_updates_status(self):
book = self.create_book() book = self.create_book()
reading = self.call("PATCH", f"/api/books/{book['id']}", {"current_page": 100}).get_json() reading = self.call(
"PATCH", f"/api/books/{book['id']}", {"current_page": 100}
).get_json()
self.assertEqual((reading["current_page"], reading["status"]), (100, "reading")) self.assertEqual((reading["current_page"], reading["status"]), (100, "reading"))
finished = self.call("PATCH", f"/api/books/{book['id']}", {"current_page": 412}).get_json() finished = self.call(
"PATCH", f"/api/books/{book['id']}", {"current_page": 412}
).get_json()
self.assertEqual(finished["status"], "finished") self.assertEqual(finished["status"], "finished")
self.assertEqual(self.call("GET", f"/api/books/{book['id']}").get_json()["current_page"], 412) self.assertEqual(
self.call("GET", f"/api/books/{book['id']}").get_json()["current_page"], 412
)
def test_edit_details(self): def test_edit_details(self):
book = self.create_book() book = self.create_book()
response = self.call( response = self.call(
"PATCH", f"/api/books/{book['id']}", "PATCH",
{"title": "Dune Messiah", "genre_id": self.genre_id("Fantasy"), "total_pages": 256}, f"/api/books/{book['id']}",
{
"title": "Dune Messiah",
"genre_id": self.genre_id("Fantasy"),
"total_pages": 256,
},
) )
self.assertEqual(response.status_code, 200) self.assertEqual(response.status_code, 200)
updated = response.get_json() updated = response.get_json()
self.assertEqual((updated["title"], updated["genre"]["name"], updated["total_pages"]), ("Dune Messiah", "Fantasy", 256)) self.assertEqual(
(updated["title"], updated["genre"]["name"], updated["total_pages"]),
("Dune Messiah", "Fantasy", 256),
)
self.assertGreater(updated["updated_at"], book["updated_at"]) self.assertGreater(updated["updated_at"], book["updated_at"])
def test_create_validation(self): def test_create_validation(self):
valid = {"title": "Dune", "author": "Frank Herbert", "genre_id": self.genre_id("Fiction"), "total_pages": 412} valid = {
"title": "Dune",
"author": "Frank Herbert",
"genre_id": self.genre_id("Fiction"),
"total_pages": 412,
}
cases = [ cases = [
({**valid, "title": " "}, "title"), ({**valid, "title": " "}, "title"),
({k: v for k, v in valid.items() if k != "author"}, "author"), ({k: v for k, v in valid.items() if k != "author"}, "author"),
@@ -69,13 +94,19 @@ class BookTests(ApiTestCase):
def test_empty_patch_is_rejected(self): def test_empty_patch_is_rejected(self):
book = self.create_book() book = self.create_book()
self.assertEqual(self.call("PATCH", f"/api/books/{book['id']}", {}).status_code, 400) self.assertEqual(
self.call("PATCH", f"/api/books/{book['id']}", {}).status_code, 400
)
def test_delete(self): def test_delete(self):
book = self.create_book() book = self.create_book()
self.assertEqual(self.call("DELETE", f"/api/books/{book['id']}").status_code, 204) self.assertEqual(
self.call("DELETE", f"/api/books/{book['id']}").status_code, 204
)
self.assertEqual(self.call("GET", f"/api/books/{book['id']}").status_code, 404) self.assertEqual(self.call("GET", f"/api/books/{book['id']}").status_code, 404)
self.assertEqual(self.call("DELETE", f"/api/books/{book['id']}").status_code, 404) self.assertEqual(
self.call("DELETE", f"/api/books/{book['id']}").status_code, 404
)
def test_out_of_range_id_is_404(self): def test_out_of_range_id_is_404(self):
self.assertEqual(self.call("GET", "/api/books/99999999999").status_code, 404) self.assertEqual(self.call("GET", "/api/books/99999999999").status_code, 404)
@@ -105,7 +136,11 @@ class BookIsolationTests(ApiTestCase):
book = self.create_book() book = self.create_book()
bob = self.other_user("bob") bob = self.other_user("bob")
path = f"/api/books/{book['id']}" path = f"/api/books/{book['id']}"
for method, payload in (("GET", None), ("PATCH", {"title": "Hacked"}), ("DELETE", None)): for method, payload in (
("GET", None),
("PATCH", {"title": "Hacked"}),
("DELETE", None),
):
with self.subTest(method=method): with self.subTest(method=method):
response = self.call(method, path, payload, client=bob) response = self.call(method, path, payload, client=bob)
self.assertEqual(response.status_code, 404) self.assertEqual(response.status_code, 404)
+25 -7
View File
@@ -29,19 +29,30 @@ class NoteTests(ApiTestCase):
def test_delete(self): def test_delete(self):
note = self.add("Temporary").get_json() note = self.add("Temporary").get_json()
self.assertEqual(self.call("DELETE", f"/api/notes/{note['id']}").status_code, 204) self.assertEqual(
self.call("DELETE", f"/api/notes/{note['id']}").status_code, 204
)
self.assertEqual(self.call("GET", self.notes_path).get_json(), []) self.assertEqual(self.call("GET", self.notes_path).get_json(), [])
self.assertEqual(self.call("DELETE", f"/api/notes/{note['id']}").status_code, 404) self.assertEqual(
self.call("DELETE", f"/api/notes/{note['id']}").status_code, 404
)
def test_validation(self): def test_validation(self):
for body, status in ((" ", 400), ("x" * 10_001, 400), ("a\x00b", 400)): for body, status in ((" ", 400), ("x" * 10_001, 400), ("a\x00b", 400)):
with self.subTest(length=len(body)): with self.subTest(length=len(body)):
self.assertEqual(self.add(body).status_code, status) self.assertEqual(self.add(body).status_code, status)
self.assertEqual(self.call("POST", self.notes_path, {"body": "ok", "book_id": 2}).status_code, 400) self.assertEqual(
self.call(
"POST", self.notes_path, {"body": "ok", "book_id": 2}
).status_code,
400,
)
def test_missing_book_and_out_of_range_ids_are_404(self): def test_missing_book_and_out_of_range_ids_are_404(self):
self.assertEqual(self.call("GET", "/api/books/999/notes").status_code, 404) self.assertEqual(self.call("GET", "/api/books/999/notes").status_code, 404)
self.assertEqual(self.call("PATCH", "/api/notes/99999999999", {"body": "x"}).status_code, 404) self.assertEqual(
self.call("PATCH", "/api/notes/99999999999", {"body": "x"}).status_code, 404
)
def test_deleting_book_deletes_its_notes(self): def test_deleting_book_deletes_its_notes(self):
self.add("Will be gone") self.add("Will be gone")
@@ -53,7 +64,9 @@ class NoteIsolationTests(ApiTestCase):
def test_other_user_cannot_read_add_edit_or_delete_notes(self): def test_other_user_cannot_read_add_edit_or_delete_notes(self):
self.register("alice") self.register("alice")
book = self.create_book() book = self.create_book()
note = self.call("POST", f"/api/books/{book['id']}/notes", {"body": "Private"}).get_json() note = self.call(
"POST", f"/api/books/{book['id']}/notes", {"body": "Private"}
).get_json()
bob = self.other_user("bob") bob = self.other_user("bob")
attempts = ( attempts = (
("GET", f"/api/books/{book['id']}/notes", None), ("GET", f"/api/books/{book['id']}/notes", None),
@@ -63,6 +76,11 @@ class NoteIsolationTests(ApiTestCase):
) )
for method, path, payload in attempts: for method, path, payload in attempts:
with self.subTest(method=method, path=path): with self.subTest(method=method, path=path):
self.assertEqual(self.call(method, path, payload, client=bob).status_code, 404) self.assertEqual(
bodies = [n["body"] for n in self.call("GET", f"/api/books/{book['id']}/notes").get_json()] self.call(method, path, payload, client=bob).status_code, 404
)
bodies = [
n["body"]
for n in self.call("GET", f"/api/books/{book['id']}/notes").get_json()
]
self.assertEqual(bodies, ["Private"]) self.assertEqual(bodies, ["Private"])
+6 -2
View File
@@ -14,7 +14,9 @@ KNOWN_ANSWER = "pbkdf2_sha256$1000$AAAAAAAAAAAAAAAAAAAAAA==$9ZLLEusnEPU3Km8h+vnd
class PasswordHashTests(unittest.TestCase): class PasswordHashTests(unittest.TestCase):
def test_known_answer(self): def test_known_answer(self):
stored = hash_password("correct horse battery staple", PEPPER, salt=bytes(16), iterations=1000) stored = hash_password(
"correct horse battery staple", PEPPER, salt=bytes(16), iterations=1000
)
self.assertEqual(stored, KNOWN_ANSWER) self.assertEqual(stored, KNOWN_ANSWER)
def test_round_trip_uses_current_iterations_and_random_salt(self): def test_round_trip_uses_current_iterations_and_random_salt(self):
@@ -28,7 +30,9 @@ class PasswordHashTests(unittest.TestCase):
self.assertFalse(verify_password("wrong password!!", KNOWN_ANSWER, PEPPER)) self.assertFalse(verify_password("wrong password!!", KNOWN_ANSWER, PEPPER))
def test_wrong_pepper_fails(self): def test_wrong_pepper_fails(self):
self.assertFalse(verify_password("correct horse battery staple", KNOWN_ANSWER, b"q" * 32)) self.assertFalse(
verify_password("correct horse battery staple", KNOWN_ANSWER, b"q" * 32)
)
def test_needs_rehash_below_current_iterations(self): def test_needs_rehash_below_current_iterations(self):
self.assertTrue(needs_rehash(KNOWN_ANSWER)) self.assertTrue(needs_rehash(KNOWN_ANSWER))
+17 -5
View File
@@ -8,7 +8,9 @@ class SearchTests(ApiTestCase):
self.scifi = self.genre_id("Science Fiction") self.scifi = self.genre_id("Science Fiction")
self.fantasy = self.genre_id("Fantasy") self.fantasy = self.genre_id("Fantasy")
self.create_book(title="Dune", author="Frank Herbert", genre_id=self.scifi) self.create_book(title="Dune", author="Frank Herbert", genre_id=self.scifi)
self.create_book(title="The Hobbit", author="J.R.R. Tolkien", genre_id=self.fantasy) self.create_book(
title="The Hobbit", author="J.R.R. Tolkien", genre_id=self.fantasy
)
self.create_book(title="100% Pure", author="Jane_Doe", genre_id=self.fantasy) self.create_book(title="100% Pure", author="Jane_Doe", genre_id=self.fantasy)
self.create_book(title="1000 Pages", author="Back\\slash", genre_id=self.scifi) self.create_book(title="1000 Pages", author="Back\\slash", genre_id=self.scifi)
@@ -24,18 +26,26 @@ class SearchTests(ApiTestCase):
def test_wildcards_match_literally(self): def test_wildcards_match_literally(self):
self.assertEqual(self.titles("q=100%25"), ["100% Pure"]) # %25 is "%" self.assertEqual(self.titles("q=100%25"), ["100% Pure"]) # %25 is "%"
self.assertEqual(self.titles("q=e_D"), ["100% Pure"]) # literal "_" in Jane_Doe self.assertEqual(self.titles("q=e_D"), ["100% Pure"]) # literal "_" in Jane_Doe
self.assertEqual(self.titles("q=k_s"), []) # unescaped "_" would match "Back\\slash" self.assertEqual(
self.titles("q=k_s"), []
) # unescaped "_" would match "Back\\slash"
self.assertEqual(self.titles("q=k%5Cs"), ["1000 Pages"]) # %5C is "\" self.assertEqual(self.titles("q=k%5Cs"), ["1000 Pages"]) # %5C is "\"
def test_genre_filter_and_combination(self): def test_genre_filter_and_combination(self):
self.assertEqual(self.titles(f"genre_id={self.fantasy}"), ["100% Pure", "The Hobbit"]) self.assertEqual(
self.titles(f"genre_id={self.fantasy}"), ["100% Pure", "The Hobbit"]
)
self.assertEqual(self.titles(f"q=100&genre_id={self.scifi}"), ["1000 Pages"]) self.assertEqual(self.titles(f"q=100&genre_id={self.scifi}"), ["1000 Pages"])
def test_blank_query_returns_everything(self): def test_blank_query_returns_everything(self):
self.assertEqual(len(self.titles("q=%20%20")), 4) self.assertEqual(len(self.titles("q=%20%20")), 4)
def test_invalid_parameters(self): def test_invalid_parameters(self):
for query, field in (("genre_id=abc", "genre_id"), ("genre_id=99999", "genre_id"), ("q=a%00b", "q")): for query, field in (
("genre_id=abc", "genre_id"),
("genre_id=99999", "genre_id"),
("q=a%00b", "q"),
):
with self.subTest(query=query): with self.subTest(query=query):
response = self.call("GET", f"/api/books?{query}") response = self.call("GET", f"/api/books?{query}")
self.assertEqual(response.status_code, 400) self.assertEqual(response.status_code, 400)
@@ -43,4 +53,6 @@ class SearchTests(ApiTestCase):
def test_search_never_returns_other_users_books(self): def test_search_never_returns_other_users_books(self):
bob = self.other_user("bob") bob = self.other_user("bob")
self.assertEqual(self.call("GET", "/api/books?q=dune", client=bob).get_json(), []) self.assertEqual(
self.call("GET", "/api/books?q=dune", client=bob).get_json(), []
)
+6 -1
View File
@@ -1,4 +1,5 @@
"""Request validation shared by every slice. Errors carry a status, a fix-it message, and the field.""" """Request validation shared by every slice. Errors carry a status, a fix-it message, and the field."""
from flask import request from flask import request
@@ -28,7 +29,11 @@ def require_utf8(value: str, field: str) -> None:
try: try:
value.encode("utf-8") value.encode("utf-8")
except UnicodeEncodeError: except UnicodeEncodeError:
raise ApiError(400, f"{field} contains unpaired surrogate characters; send valid UTF-8 text", field) from None raise ApiError(
400,
f"{field} contains unpaired surrogate characters; send valid UTF-8 text",
field,
) from None
def text(body: dict, field: str, max_len: int) -> str: def text(body: dict, field: str, max_len: int) -> str: