From 63c99914210f389c10cc9892e3aea7c3d075d048 Mon Sep 17 00:00:00 2001 From: Malcolm Roberts Date: Fri, 2 Oct 2026 14:46:16 -0500 Subject: [PATCH] style: format backend with black Also drops the session-sliding comment in login_required. --- backend/app.py | 24 ++++++++-- backend/auth.py | 56 +++++++++++++++++------ backend/books.py | 53 +++++++++++++++++---- backend/db.py | 5 +- backend/notes.py | 13 +++++- backend/tests/__init__.py | 13 ++++-- backend/tests/support.py | 7 ++- backend/tests/test_app.py | 14 ++++-- backend/tests/test_auth.py | 81 +++++++++++++++++++++++++++------ backend/tests/test_books.py | 61 +++++++++++++++++++------ backend/tests/test_notes.py | 32 ++++++++++--- backend/tests/test_passwords.py | 8 +++- backend/tests/test_search.py | 28 ++++++++---- backend/validation.py | 7 ++- 14 files changed, 320 insertions(+), 82 deletions(-) diff --git a/backend/app.py b/backend/app.py index e05ae34..3216aec 100644 --- a/backend/app.py +++ b/backend/app.py @@ -1,4 +1,5 @@ """Application entry point: wires slices, cross-cutting request rules, and JSON error handling.""" + import logging import psycopg2.errors @@ -16,7 +17,9 @@ MUTATING_METHODS = {"POST", "PUT", "PATCH", "DELETE"} 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__) # 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). @@ -35,11 +38,19 @@ def create_app() -> Flask: # CSRF defense (ADR-0002): cross-site requests can only send this content type after a # CORS preflight, and this API never grants CORS. 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 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 @app.errorhandler(ApiError) @@ -55,7 +66,12 @@ def create_app() -> Flask: @app.errorhandler(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) def handle_unexpected(_error: Exception): diff --git a/backend/auth.py b/backend/auth.py index 5f9a34c..8533889 100644 --- a/backend/auth.py +++ b/backend/auth.py @@ -1,4 +1,5 @@ """Auth slice: password hashing (ADR-0001), server-side sessions (ADR-0002), and auth routes.""" + import base64 import functools import hashlib @@ -35,7 +36,14 @@ def init_app(app: Flask) -> None: # --- 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 derived = _derive(password, pepper, salt, iterations) return f"pbkdf2_sha256${iterations}${_b64(salt)}${_b64(derived)}" @@ -98,14 +106,13 @@ def load_pepper(pepper_file: Path) -> bytes: # --- 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 @@ -135,25 +142,32 @@ def _pepper() -> bytes: 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( + "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") + 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") + raise ApiError( + 400, f"password must be at most {MAX_PASSWORD} characters", "password" + ) require_utf8(password, "password") return password @@ -163,17 +177,25 @@ 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") + 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") + 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 + raise ApiError( + 409, "Username already taken; choose another", "username" + ) from None return _start_session(user, 201) @@ -185,7 +207,8 @@ def login(): 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,) + "SELECT id, username, password_hash FROM users WHERE lower(username) = lower(%s)", + (username,), ) if user is None: verify_password(password, _dummy_hash(), _pepper()) @@ -193,7 +216,10 @@ def 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"])) + query( + "UPDATE users SET password_hash = %s WHERE id = %s", + (hash_password(password, _pepper()), user["id"]), + ) return _start_session(user, 200) @@ -203,11 +229,15 @@ def logout(): 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") + 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,))) + return jsonify( + query_one("SELECT id, username FROM users WHERE id = %s", (g.user_id,)) + ) diff --git a/backend/books.py b/backend/books.py index bd82de6..0fe0afb 100644 --- a/backend/books.py +++ b/backend/books.py @@ -1,4 +1,5 @@ """Books slice: CRUD, progress, search/filter, and the genre list.""" + import re import psycopg2.errors @@ -60,7 +61,9 @@ def escape_like(term: str) -> str: 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: raise ApiError(404, "Book not found") return row @@ -80,7 +83,9 @@ def _save(sql: str, params: tuple) -> dict | None: try: return query_one(sql, params) 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") @@ -97,14 +102,18 @@ def list_books(): search = request.args.get("q", "").strip() if 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)}%" sql += " AND (b.title ILIKE %s OR b.author ILIKE %s)" params += [pattern, pattern] genre_id = request.args.get("genre_id", "") if 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" params.append(int(genre_id)) sql += " ORDER BY b.updated_at DESC, b.id DESC" @@ -116,7 +125,9 @@ def list_books(): def create_book(): body = json_body(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") row = _save( """ @@ -124,7 +135,14 @@ def create_book(): VALUES (%s, %s, %s, %s, %s, %s) 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 @@ -143,7 +161,9 @@ def update_book(book_id: int): if not body: raise ApiError(400, f"Provide at least one of: {', '.join(sorted(FIELDS))}") 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") _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() 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))) @@ -160,6 +187,12 @@ def update_book(book_id: int): @bp.delete("/books/") @login_required 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") return "", 204 diff --git a/backend/db.py b/backend/db.py index 1798462..aa4698c 100644 --- a/backend/db.py +++ b/backend/db.py @@ -1,4 +1,5 @@ """PostgreSQL access: a connection pool, one connection per request, and two query helpers.""" + import os from pathlib import Path @@ -11,7 +12,9 @@ SCHEMA = Path(__file__).with_name("schema.sql") 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() try: with conn, conn.cursor() as cur: diff --git a/backend/notes.py b/backend/notes.py index 3714bd9..d331d8d 100644 --- a/backend/notes.py +++ b/backend/notes.py @@ -1,4 +1,5 @@ """Notes slice: a per-book reading journal. Ownership is always derived through books.user_id.""" + from flask import Blueprint, g, jsonify from auth import login_required @@ -22,7 +23,12 @@ def to_json(row: dict) -> dict: 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") @@ -42,7 +48,10 @@ def list_notes(book_id: int): def add_note(book_id: int): _require_book(book_id) 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 diff --git a/backend/tests/__init__.py b/backend/tests/__init__.py index 476ae89..dbae783 100644 --- a/backend/tests/__init__.py +++ b/backend/tests/__init__.py @@ -1,10 +1,13 @@ """Point the app at a dedicated test database before any test module imports it.""" + import os import psycopg2 from psycopg2 import sql -ADMIN_URL = os.environ.get("TEST_ADMIN_DATABASE_URL", "postgresql://postgres:postgres@127.0.0.1:5432/postgres") +ADMIN_URL = os.environ.get( + "TEST_ADMIN_DATABASE_URL", "postgresql://postgres:postgres@127.0.0.1:5432/postgres" +) TEST_DATABASE = "books_test" @@ -13,9 +16,13 @@ def _ensure_test_database() -> None: conn.autocommit = True # CREATE DATABASE cannot run inside a transaction try: 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: - cur.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(TEST_DATABASE))) + cur.execute( + sql.SQL("CREATE DATABASE {}").format(sql.Identifier(TEST_DATABASE)) + ) finally: conn.close() diff --git a/backend/tests/support.py b/backend/tests/support.py index c686999..5c5983c 100644 --- a/backend/tests/support.py +++ b/backend/tests/support.py @@ -34,7 +34,12 @@ class ApiTestCase(unittest.TestCase): return (client or self.client).open(path, **kwargs) 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()) return response diff --git a/backend/tests/test_app.py b/backend/tests/test_app.py index b8b2945..e91ba8f 100644 --- a/backend/tests/test_app.py +++ b/backend/tests/test_app.py @@ -16,7 +16,9 @@ class AppTests(ApiTestCase): 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 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): response = self.call("GET", "/api/nope") @@ -24,7 +26,9 @@ class AppTests(ApiTestCase): self.assertIn("error", response.get_json()) 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.get_json(), @@ -32,6 +36,10 @@ class AppTests(ApiTestCase): ) 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.assertIn("error", response.get_json()) diff --git a/backend/tests/test_auth.py b/backend/tests/test_auth.py index eb89645..df87fbf 100644 --- a/backend/tests/test_auth.py +++ b/backend/tests/test_auth.py @@ -8,7 +8,11 @@ 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 + 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") @@ -25,11 +29,20 @@ class AuthTests(ApiTestCase): 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) + 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"}) + 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") @@ -40,7 +53,14 @@ class AuthTests(ApiTestCase): ({"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"), + ( + { + "username": "alice", + "password": "long enough password", + "is_admin": True, + }, + "is_admin", + ), ] for payload, field in cases: with self.subTest(payload=payload): @@ -51,33 +71,54 @@ class AuthTests(ApiTestCase): 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"}) + 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"}) + 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"}) + 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"}) + 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.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): @@ -87,14 +128,18 @@ class AuthTests(ApiTestCase): 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'") + 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] + 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): @@ -102,6 +147,14 @@ class AuthTests(ApiTestCase): 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"}) + 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$")) + self.assertTrue( + db_execute("SELECT password_hash FROM users")[0][0].startswith( + "pbkdf2_sha256$600000$" + ) + ) diff --git a/backend/tests/test_books.py b/backend/tests/test_books.py index 5d892c7..380349f 100644 --- a/backend/tests/test_books.py +++ b/backend/tests/test_books.py @@ -9,8 +9,14 @@ class BookTests(ApiTestCase): def test_create_returns_full_shape(self): book = self.create_book() self.assertEqual(book["title"], "Dune") - self.assertEqual(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.assertEqual( + 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 def test_create_trims_text(self): @@ -19,25 +25,44 @@ class BookTests(ApiTestCase): def test_progress_updates_status(self): 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")) - 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(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): book = self.create_book() response = self.call( - "PATCH", f"/api/books/{book['id']}", - {"title": "Dune Messiah", "genre_id": self.genre_id("Fantasy"), "total_pages": 256}, + "PATCH", + f"/api/books/{book['id']}", + { + "title": "Dune Messiah", + "genre_id": self.genre_id("Fantasy"), + "total_pages": 256, + }, ) self.assertEqual(response.status_code, 200) 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"]) 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 = [ ({**valid, "title": " "}, "title"), ({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): 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): 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("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): self.assertEqual(self.call("GET", "/api/books/99999999999").status_code, 404) @@ -105,7 +136,11 @@ class BookIsolationTests(ApiTestCase): book = self.create_book() bob = self.other_user("bob") 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): response = self.call(method, path, payload, client=bob) self.assertEqual(response.status_code, 404) diff --git a/backend/tests/test_notes.py b/backend/tests/test_notes.py index d6c3b47..bb21f5b 100644 --- a/backend/tests/test_notes.py +++ b/backend/tests/test_notes.py @@ -29,19 +29,30 @@ class NoteTests(ApiTestCase): def test_delete(self): 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("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): for body, status in ((" ", 400), ("x" * 10_001, 400), ("a\x00b", 400)): with self.subTest(length=len(body)): 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): 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): self.add("Will be gone") @@ -53,7 +64,9 @@ class NoteIsolationTests(ApiTestCase): def test_other_user_cannot_read_add_edit_or_delete_notes(self): self.register("alice") 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") attempts = ( ("GET", f"/api/books/{book['id']}/notes", None), @@ -63,6 +76,11 @@ class NoteIsolationTests(ApiTestCase): ) for method, path, payload in attempts: with self.subTest(method=method, path=path): - self.assertEqual(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( + 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"]) diff --git a/backend/tests/test_passwords.py b/backend/tests/test_passwords.py index f386613..a406675 100644 --- a/backend/tests/test_passwords.py +++ b/backend/tests/test_passwords.py @@ -14,7 +14,9 @@ KNOWN_ANSWER = "pbkdf2_sha256$1000$AAAAAAAAAAAAAAAAAAAAAA==$9ZLLEusnEPU3Km8h+vnd class PasswordHashTests(unittest.TestCase): 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) 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)) 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): self.assertTrue(needs_rehash(KNOWN_ANSWER)) diff --git a/backend/tests/test_search.py b/backend/tests/test_search.py index b078e0c..a2ee0ff 100644 --- a/backend/tests/test_search.py +++ b/backend/tests/test_search.py @@ -8,7 +8,9 @@ class SearchTests(ApiTestCase): self.scifi = self.genre_id("Science Fiction") self.fantasy = self.genre_id("Fantasy") 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="1000 Pages", author="Back\\slash", genre_id=self.scifi) @@ -22,20 +24,28 @@ class SearchTests(ApiTestCase): self.assertEqual(self.titles("q=tolkien"), ["The Hobbit"]) def test_wildcards_match_literally(self): - 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=k_s"), []) # unescaped "_" would match "Back\\slash" - self.assertEqual(self.titles("q=k%5Cs"), ["1000 Pages"]) # %5C 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=k_s"), [] + ) # unescaped "_" would match "Back\\slash" + self.assertEqual(self.titles("q=k%5Cs"), ["1000 Pages"]) # %5C is "\" 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"]) def test_blank_query_returns_everything(self): self.assertEqual(len(self.titles("q=%20%20")), 4) 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): response = self.call("GET", f"/api/books?{query}") self.assertEqual(response.status_code, 400) @@ -43,4 +53,6 @@ class SearchTests(ApiTestCase): def test_search_never_returns_other_users_books(self): 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(), [] + ) diff --git a/backend/validation.py b/backend/validation.py index b487728..e0d0da4 100644 --- a/backend/validation.py +++ b/backend/validation.py @@ -1,4 +1,5 @@ """Request validation shared by every slice. Errors carry a status, a fix-it message, and the field.""" + from flask import request @@ -28,7 +29,11 @@ def require_utf8(value: str, field: str) -> None: try: value.encode("utf-8") 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: