feat: scaffold Flask + React app with schema and JSON error handling

Claude-Session: https://claude.ai/code/session_01M9MLit5Ko3X4s7rzC5Kv7X
This commit is contained in:
2026-10-02 14:11:37 -05:00
parent 02d4598697
commit b15448f109
23 changed files with 4275 additions and 0 deletions
+68
View File
@@ -0,0 +1,68 @@
"""Application entry point: wires slices, cross-cutting request rules, and JSON error handling."""
import logging
import psycopg2.errors
from flask import Flask, g, jsonify, request
from werkzeug.exceptions import HTTPException
import auth
import books
import db
import notes
from validation import ApiError
log = logging.getLogger(__name__)
MUTATING_METHODS = {"POST", "PUT", "PATCH", "DELETE"}
def create_app() -> Flask:
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")
app = Flask(__name__)
db.init_app(app)
auth.init_app(app)
app.register_blueprint(books.bp)
app.register_blueprint(notes.bp)
@app.get("/api/health")
def health():
return {"status": "ok"}
@app.before_request
def require_json_for_mutations():
# 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")
@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"))
return response
@app.errorhandler(ApiError)
def handle_api_error(error: ApiError):
body = {"error": error.message}
if error.field:
body["field"] = error.field
return jsonify(body), error.status
@app.errorhandler(HTTPException)
def handle_http_error(error: HTTPException):
return jsonify(error=error.description), error.code
@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
@app.errorhandler(Exception)
def handle_unexpected(_error: Exception):
log.exception("Unhandled error on %s %s", request.method, request.path)
return jsonify(error="Unexpected server error"), 500
return app
app = create_app()
if __name__ == "__main__":
app.run(host="127.0.0.1", port=5000)
+8
View File
@@ -0,0 +1,8 @@
"""Auth slice: password hashing (ADR-0001), server-side sessions (ADR-0002), and auth routes."""
from flask import Blueprint, Flask
bp = Blueprint("auth", __name__, url_prefix="/api/auth")
def init_app(app: Flask) -> None:
app.register_blueprint(bp)
+4
View File
@@ -0,0 +1,4 @@
"""Books slice: CRUD, progress, search/filter, and the genre list."""
from flask import Blueprint
bp = Blueprint("books", __name__, url_prefix="/api")
+58
View File
@@ -0,0 +1,58 @@
"""PostgreSQL access: a connection pool, one connection per request, and two query helpers."""
import os
from pathlib import Path
import psycopg2.extras
import psycopg2.pool
from flask import Flask, current_app, g
DEFAULT_DATABASE_URL = "postgresql://postgres:[email protected]:5432/postgres"
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))
conn = pool.getconn()
try:
with conn, conn.cursor() as cur:
cur.execute(SCHEMA.read_text())
finally:
pool.putconn(conn)
app.extensions["db_pool"] = pool
app.after_request(_finish_transaction)
app.teardown_appcontext(_release_connection)
def query(sql: str, params: tuple = ()) -> list[dict]:
with _connection().cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
cur.execute(sql, params)
return cur.fetchall() if cur.description else []
def query_one(sql: str, params: tuple = ()) -> dict | None:
rows = query(sql, params)
return rows[0] if rows else None
def _connection():
if "db" not in g:
g.db = current_app.extensions["db_pool"].getconn()
return g.db
def _finish_transaction(response):
# Commit only successful responses so a 4xx/5xx never leaves partial writes behind.
conn = g.get("db")
if conn is not None:
if response.status_code < 400:
conn.commit()
else:
conn.rollback()
return response
def _release_connection(_exc):
conn = g.pop("db", None)
if conn is not None:
conn.rollback() # no-op after commit; discards work if after_request never ran
current_app.extensions["db_pool"].putconn(conn)
+4
View File
@@ -0,0 +1,4 @@
"""Notes slice: a per-book reading journal."""
from flask import Blueprint
bp = Blueprint("notes", __name__, url_prefix="/api")
+2
View File
@@ -0,0 +1,2 @@
Flask>=3.0.0
psycopg2-binary>=2.9.9
+52
View File
@@ -0,0 +1,52 @@
CREATE TABLE IF NOT EXISTS users (
id SERIAL PRIMARY KEY,
username TEXT NOT NULL,
password_hash TEXT NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
CREATE UNIQUE INDEX IF NOT EXISTS users_username_lower_key ON users (lower(username));
CREATE TABLE IF NOT EXISTS sessions (
token_hash BYTEA PRIMARY KEY,
user_id INT NOT NULL REFERENCES users ON DELETE CASCADE,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
expires_at TIMESTAMPTZ NOT NULL
);
CREATE INDEX IF NOT EXISTS sessions_user_id_idx ON sessions (user_id);
CREATE TABLE IF NOT EXISTS genres (
id SMALLSERIAL PRIMARY KEY,
name TEXT NOT NULL UNIQUE
);
-- WHERE NOT EXISTS instead of ON CONFLICT: ON CONFLICT consumes a sequence value per row
-- on every startup, which would eventually overflow SMALLSERIAL.
INSERT INTO genres (name)
SELECT seed.name
FROM (VALUES ('Fiction'), ('Non-Fiction'), ('Mystery'), ('Thriller'), ('Science Fiction'),
('Fantasy'), ('Romance'), ('Horror'), ('Historical Fiction'), ('Biography'),
('Memoir'), ('Self-Help'), ('Science'), ('History'), ('Poetry'), ('Other')) AS seed(name)
WHERE NOT EXISTS (SELECT 1 FROM genres WHERE genres.name = seed.name);
CREATE TABLE IF NOT EXISTS books (
id SERIAL PRIMARY KEY,
user_id INT NOT NULL REFERENCES users ON DELETE CASCADE,
title TEXT NOT NULL CHECK (btrim(title) <> ''),
author TEXT NOT NULL CHECK (btrim(author) <> ''),
genre_id SMALLINT NOT NULL REFERENCES genres,
total_pages INT NOT NULL CHECK (total_pages > 0),
current_page INT NOT NULL DEFAULT 0,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
CONSTRAINT books_current_page_range CHECK (current_page BETWEEN 0 AND total_pages)
);
CREATE INDEX IF NOT EXISTS books_user_id_idx ON books (user_id);
CREATE TABLE IF NOT EXISTS notes (
id SERIAL PRIMARY KEY,
book_id INT NOT NULL REFERENCES books ON DELETE CASCADE,
body TEXT NOT NULL CHECK (btrim(body) <> ''),
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
CREATE INDEX IF NOT EXISTS notes_book_id_idx ON notes (book_id);
+26
View File
@@ -0,0 +1,26 @@
"""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:[email protected]:5432/postgres")
TEST_DATABASE = "books_test"
def _ensure_test_database() -> None:
conn = psycopg2.connect(ADMIN_URL)
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,))
if cur.fetchone() is None:
cur.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(TEST_DATABASE)))
finally:
conn.close()
_ensure_test_database()
# Assigned, never defaulted: tests TRUNCATE tables, so they must not inherit a real DATABASE_URL.
os.environ["DATABASE_URL"] = ADMIN_URL.rsplit("/", 1)[0] + "/" + TEST_DATABASE
os.environ["PASSWORD_PEPPER"] = "test-pepper-0123456789abcdef0123456789abcdef"
+58
View File
@@ -0,0 +1,58 @@
import os
import unittest
import psycopg2
from app import app
HTTPS = "https://localhost" # the session cookie is Secure
def db_execute(sql: str, params: tuple = ()) -> list[tuple]:
conn = psycopg2.connect(os.environ["DATABASE_URL"])
try:
with conn, conn.cursor() as cur:
cur.execute(sql, params)
return cur.fetchall() if cur.description else []
finally:
conn.close()
class ApiTestCase(unittest.TestCase):
def setUp(self):
if not os.environ["DATABASE_URL"].endswith("/books_test"):
raise RuntimeError("Refusing to truncate a database that is not books_test")
db_execute("TRUNCATE users, sessions, books, notes RESTART IDENTITY CASCADE")
self.client = app.test_client()
def call(self, method, path, json=None, client=None):
kwargs = {"method": method, "base_url": HTTPS}
if json is not None:
kwargs["json"] = json
elif method != "GET":
kwargs["content_type"] = "application/json"
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)
self.assertEqual(response.status_code, 201, response.get_json())
return response
def other_user(self, username="bob"):
client = app.test_client()
self.register(username, client=client)
return client
def genre_id(self, name: str) -> int:
return db_execute("SELECT id FROM genres WHERE name = %s", (name,))[0][0]
def create_book(self, client=None, **overrides) -> dict:
payload = {
"title": "Dune",
"author": "Frank Herbert",
"genre_id": self.genre_id("Science Fiction"),
"total_pages": 412,
} | overrides
response = self.call("POST", "/api/books", payload, client)
self.assertEqual(response.status_code, 201, response.get_json())
return response.get_json()
+32
View File
@@ -0,0 +1,32 @@
from pathlib import Path
from tests.support import HTTPS, ApiTestCase, db_execute
SCHEMA = Path(__file__).parent.parent / "schema.sql"
class AppTests(ApiTestCase):
def test_health(self):
response = self.call("GET", "/api/health")
self.assertEqual(response.status_code, 200)
self.assertEqual(response.get_json(), {"status": "ok"})
def test_genres_seeded_and_reseeding_is_idempotent(self):
before = db_execute("SELECT last_value FROM genres_id_seq")[0][0]
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)
def test_unknown_route_returns_json_404(self):
response = self.call("GET", "/api/nope")
self.assertEqual(response.status_code, 404)
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)
self.assertEqual(response.status_code, 400)
self.assertEqual(
response.get_json(),
{"error": "Request body must be JSON with Content-Type: application/json"},
)
+54
View File
@@ -0,0 +1,54 @@
"""Request validation shared by every slice. Errors carry a status, a fix-it message, and the field."""
from flask import request
class ApiError(Exception):
def __init__(self, status: int, message: str, field: str | None = None):
super().__init__(message)
self.status = status
self.message = message
self.field = field
def json_body(allowed: set[str]) -> dict:
body = request.get_json(silent=True)
if not isinstance(body, dict):
raise ApiError(400, "Request body must be a JSON object")
unknown = sorted(set(body) - allowed)
if unknown:
raise ApiError(
400,
f"Unknown field(s): {', '.join(unknown)}. Allowed: {', '.join(sorted(allowed))}",
unknown[0],
)
return body
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
def text(body: dict, field: str, max_len: int) -> str:
value = body.get(field)
if not isinstance(value, str) or not value.strip():
raise ApiError(400, f"{field} is required and must be non-blank text", field)
value = value.strip()
if len(value) > max_len:
raise ApiError(400, f"{field} must be at most {max_len} characters", field)
if "\x00" in value:
raise ApiError(400, f"{field} must not contain NUL characters", field)
require_utf8(value, field)
return value
def integer(body: dict, field: str, low: int, high: int) -> int:
value = body.get(field)
# bool is a subclass of int in Python; JSON true must not count as 1.
if isinstance(value, bool) or not isinstance(value, int):
raise ApiError(400, f"{field} must be a whole number", field)
if not low <= value <= high:
raise ApiError(400, f"{field} must be between {low} and {high}", field)
return value