"""PostgreSQL access: a connection pool, one connection per request, and query helpers.""" import logging import os from collections.abc import Sequence from pathlib import Path from typing import Any import psycopg2.extensions import psycopg2.extras import psycopg2.pool from flask import Flask, Response, current_app, g log = logging.getLogger(__name__) DEFAULT_DATABASE_URL = "postgresql://postgres:postgres@127.0.0.1:5432/postgres" SCHEMA = Path(__file__).with_name("schema.sql") Row = dict[str, Any] Params = tuple[object, ...] 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: Params = ()) -> Sequence[Row]: 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: Params = ()) -> Row | None: rows = query(sql, params) return rows[0] if rows else None def query_row(sql: str, params: Params = ()) -> Row: rows = query(sql, params) if len(rows) != 1: raise RuntimeError( f"Expected exactly one row but got {len(rows)}; the statement must return one row: {sql.strip()}" ) return rows[0] def _connection() -> psycopg2.extensions.connection: if "db" not in g: g.db = current_app.extensions["db_pool"].getconn() conn: psycopg2.extensions.connection = g.db return conn def _finish_transaction(response: Response) -> Response: 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: BaseException | None) -> None: conn = g.pop("db", None) if conn is None: return pool = current_app.extensions["db_pool"] try: conn.rollback() except psycopg2.Error: log.warning( "Discarding a database connection that failed to roll back; the pool will open a new one", exc_info=True, ) pool.putconn(conn, close=True) else: pool.putconn(conn)