Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/register_your_data_api/auth/fga/fga_provider_db.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ def __init__(self, connection_str: str):
self._connection_str = connection_str

def setup(self) -> None:
self._engine = create_engine(self._connection_str, echo=True)
self._engine = create_engine(self._connection_str, echo=True, pool_pre_ping=True)

def get_user_fine_grained_permissions(self, user: UUID) -> list[FineGrainedAuthorisationRoleAssociation]:
with Session(self._engine) as session:
Expand Down
38 changes: 38 additions & 0 deletions src/register_your_data_api/exception_handlers.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import fastapi
import fastapi.responses
import sqlalchemy.exc
import starlette.exceptions
from fastapi.exceptions import RequestValidationError
from starlette.requests import Request
Expand Down Expand Up @@ -94,6 +95,41 @@ async def validation_exception_handler(
)


async def db_operational_error_handler(
request: Request, exc: sqlalchemy.exc.OperationalError
) -> fastapi.responses.JSONResponse:
"""Exception handler for database connectivity errors (e.g. the connection being dropped
or the database being unreachable). These are runtime conditions outside the application's
control that a retry is likely to resolve, so they are reported as a 503 rather than a 500.
"""

context = request.app.state.context # type: Context

context.app_logger.error(f"A database connectivity error occurred: {exc}", exc_info=True)

return fastapi.responses.JSONResponse(
{"status": "failed", "data": None, "error": {"status_code": 503, "error_msg": "Service Unavailable"}},
status_code=503,
)


async def db_error_handler(request: Request, exc: sqlalchemy.exc.DBAPIError) -> fastapi.responses.JSONResponse:
"""Exception handler for database errors that are not connectivity issues (e.g. IntegrityError,
DataError). Unlike OperationalError, retrying these will not help, so they are reported as a 500,
while still being logged clearly as database errors rather than falling through to the generic
unhandled exception handler.
"""

context = request.app.state.context # type: Context

context.app_logger.error(f"A database error occurred: {exc}", exc_info=True)

return fastapi.responses.JSONResponse(
{"status": "failed", "data": None, "error": {"status_code": 500, "error_msg": "Internal Server Error"}},
status_code=500,
)


async def unhandled_exception_handler(request: Request, exc: Exception) -> fastapi.responses.JSONResponse:
"""Catches all unhandled exceptions and returns a generic 500 server error with a simple error message"""

Expand All @@ -118,4 +154,6 @@ def add_exception_handlers(app: fastapi.FastAPI) -> None:
app.add_exception_handler(RYDUserException, ryd_user_exception_handler) # type: ignore[arg-type]
app.add_exception_handler(RequestValidationError, validation_exception_handler) # type: ignore[arg-type]
app.add_exception_handler(starlette.exceptions.HTTPException, http_exception_handler) # type: ignore[arg-type]
app.add_exception_handler(sqlalchemy.exc.OperationalError, db_operational_error_handler) # type: ignore[arg-type]
app.add_exception_handler(sqlalchemy.exc.DBAPIError, db_error_handler) # type: ignore[arg-type]
app.add_exception_handler(Exception, unhandled_exception_handler)
103 changes: 103 additions & 0 deletions tests/unit/test_exception_handlers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
"""Tests for the application's exception handlers."""

import asyncio
import json
from unittest import mock

import fastapi
import sqlalchemy.exc
from fastapi.testclient import TestClient

from register_your_data_api import exception_handlers


def test_db_operational_error_handler_logs_and_returns_service_unavailable() -> None:
"""A dropped/unavailable database connection (e.g. psycopg.errors.AdminShutdown surfacing as a
sqlalchemy.exc.OperationalError) is a runtime condition a retry is likely to resolve, so it must
be logged clearly, distinct from other unhandled errors, and reported to the client as a 503."""

exc = sqlalchemy.exc.OperationalError(
statement="SELECT 1",
params=None,
orig=Exception("terminating connection due to administrator command"),
)

request = mock.MagicMock()
context = request.app.state.context

response = asyncio.run(exception_handlers.db_operational_error_handler(request, exc))

assert response.status_code == 503
assert json.loads(bytes(response.body)) == {
"status": "failed",
"data": None,
"error": {"status_code": 503, "error_msg": "Service Unavailable"},
}

context.app_logger.error.assert_called_once_with(f"A database connectivity error occurred: {exc}", exc_info=True)


def test_db_error_handler_logs_and_returns_internal_server_error() -> None:
"""A database error that isn't a connectivity problem (e.g. IntegrityError from bad data) will
not be fixed by the client retrying, so it must be reported as a 500, while still being logged
clearly as a database error rather than falling through to the generic unhandled handler."""

exc = sqlalchemy.exc.IntegrityError(
statement="INSERT INTO foo VALUES (1)",
params=None,
orig=Exception("duplicate key value violates unique constraint"),
)

request = mock.MagicMock()
context = request.app.state.context

response = asyncio.run(exception_handlers.db_error_handler(request, exc))

assert response.status_code == 500
assert json.loads(bytes(response.body)) == {
"status": "failed",
"data": None,
"error": {"status_code": 500, "error_msg": "Internal Server Error"},
}

context.app_logger.error.assert_called_once_with(f"A database error occurred: {exc}", exc_info=True)


def test_handlers_registered_for_operational_error_and_dbapierror() -> None:
"""Both handlers must be registered: the specific OperationalError handler for connectivity
errors, and the DBAPIError handler for other database errors (e.g. IntegrityError, which is a
DBAPIError but not an OperationalError)."""

app = mock.MagicMock()

exception_handlers.add_exception_handlers(app)

app.add_exception_handler.assert_any_call(
sqlalchemy.exc.OperationalError, exception_handlers.db_operational_error_handler
)
app.add_exception_handler.assert_any_call(sqlalchemy.exc.DBAPIError, exception_handlers.db_error_handler)


def test_operational_error_and_dbapierror_use_different_handlers() -> None:
"""End-to-end check of FastAPI's exception dispatch (not just the handler bodies in isolation):
since OperationalError is itself a DBAPIError subclass, registering a handler for each only
behaves correctly if FastAPI picks the most specific match. An OperationalError must be routed
to the 503 handler, while a sibling DBAPIError subclass (IntegrityError, not fixed by retrying)
must fall through to the 500 handler rather than both being treated the same."""

app = fastapi.FastAPI()
app.state.context = mock.MagicMock()
exception_handlers.add_exception_handlers(app)

@app.get("/operational-error")
def raise_operational_error() -> None:
raise sqlalchemy.exc.OperationalError(statement="SELECT 1", params=None, orig=Exception("connection lost"))

@app.get("/integrity-error")
def raise_integrity_error() -> None:
raise sqlalchemy.exc.IntegrityError(statement="INSERT INTO foo VALUES (1)", params=None, orig=Exception("x"))

client = TestClient(app, raise_server_exceptions=False)

assert client.get("/operational-error").status_code == 503
assert client.get("/integrity-error").status_code == 500
Loading