diff --git a/app/api/django/handlers.py b/app/api/django/handlers.py new file mode 100644 index 0000000..df71acb --- /dev/null +++ b/app/api/django/handlers.py @@ -0,0 +1,5 @@ +from django.http import HttpRequest, JsonResponse + + +def custom_handler404(request: HttpRequest, exception: Exception) -> JsonResponse: + return JsonResponse(data={"detail": "Not Found"}, status=404) diff --git a/app/api/django/items/routes.py b/app/api/django/items/routes.py index d778c65..7e99593 100644 --- a/app/api/django/items/routes.py +++ b/app/api/django/items/routes.py @@ -7,6 +7,7 @@ from app.api.django.types import EnhancedHttpRequest from app.core.django.context import Context +from app.domain.dev.commands import ItemCreateError, create_item_error_command from app.domain.items.commands import ( create_item_command, delete_item_command, @@ -56,3 +57,17 @@ def delete(self, request: HttpRequest, item_id: uuid.UUID) -> JsonResponse: context = Context() delete_item_command(context, item_id=item_id) return JsonResponse(data="", status=204, safe=False) + + +class ItemViewSpecial(View): + body_models: ClassVar[dict[str, type[BaseModel]]] = {"POST": ItemCreateError} + + def post(self, request: EnhancedHttpRequest[ItemCreateError]) -> JsonResponse: + error_type = request.GET.get("error_type") + context = Context() + item = create_item_error_command( + context, + item_create=request.validated_data, + error_type=error_type, # ty:ignore[invalid-argument-type] + ) + return JsonResponse(data=item.model_dump(), safe=False) diff --git a/app/api/django/middlewares.py b/app/api/django/middlewares.py index ebd99af..126e76a 100644 --- a/app/api/django/middlewares.py +++ b/app/api/django/middlewares.py @@ -3,7 +3,7 @@ from typing import Any from django.db import transaction -from django.http import HttpRequest, HttpResponse, JsonResponse +from django.http import HttpRequest, HttpResponse from pydantic import BaseModel, ValidationError from app.api.django.types import EnhancedHttpRequest @@ -11,7 +11,6 @@ from app.domain.exceptions import DomainError DJANGO_MIDDLEWARES = [ - "app.api.django.middlewares.DomainExceptionMiddleware", "app.api.django.middlewares.PydanticValidationMiddleware", "app.api.django.middlewares.TransactionMiddleware", ] @@ -35,27 +34,6 @@ def get_json_body(self, request: HttpRequest) -> HttpResponse | dict[str, Any]: ) -class DomainExceptionMiddleware(BaseMiddleware): - def process_exception( - self, request: HttpRequest, exc: Exception - ) -> HttpResponse | None: - if not isinstance(exc, DomainError): - return None - - for error_cls in type(exc).mro(): - if issubclass(error_cls, DomainError) and error_cls in ERROR_MAPPING: - return JsonResponse( - data={"detail": str(exc)}, - status=ERROR_MAPPING[error_cls], - ) - - return HttpResponse( - content="Internal Server Error", - status=500, - content_type="text/plain", - ) - - class PydanticValidationMiddleware(BaseMiddleware): def process_view[T: BaseModel]( self, @@ -92,3 +70,36 @@ class TransactionMiddleware(BaseMiddleware): def __call__(self, request: HttpRequest) -> HttpResponse: with transaction.atomic(): return self.get_response(request) + + def process_exception( + self, + request: HttpRequest, + exc: Exception, + ) -> HttpResponse: + # TODO: `transaction.atomic()` context manager does not handle rollback on error + # is there a better way to do that? + transaction.set_rollback(True) + + if isinstance(exc, DomainError): + return handle_domain_exceptions(exc) + + return HttpResponse( + content=json.dumps({"detail": "Internal Server Error"}), + status=500, + content_type="application/json", + ) + + +def handle_domain_exceptions(exc: DomainError) -> HttpResponse: + status_code = 500 + + for error_cls in type(exc).mro(): + if issubclass(error_cls, DomainError) and error_cls in ERROR_MAPPING: + status_code = ERROR_MAPPING[error_cls] + break + + return HttpResponse( + content=json.dumps({"detail": str(exc)}), + status=status_code, + content_type="application/json", + ) diff --git a/app/api/django/urls.py b/app/api/django/urls.py index 326ccc9..9feb1f8 100644 --- a/app/api/django/urls.py +++ b/app/api/django/urls.py @@ -1,9 +1,14 @@ from django.urls import path from django.urls.resolvers import URLPattern, URLResolver -from app.api.django.items.routes import ItemView, ItemViewDetail +from app.api.django.handlers import custom_handler404 +from app.api.django.items.routes import ItemView, ItemViewDetail, ItemViewSpecial urlpatterns: list[URLPattern | URLResolver] = [ path("items", ItemView.as_view()), path("items/", ItemViewDetail.as_view()), + path("dev/error", ItemViewSpecial.as_view()), ] + +# Override handler to return JSON response +handler404 = custom_handler404 diff --git a/app/api/fastapi/app.py b/app/api/fastapi/app.py index 1f8f63c..702d76f 100644 --- a/app/api/fastapi/app.py +++ b/app/api/fastapi/app.py @@ -1,5 +1,6 @@ from fastapi import FastAPI +from app.api.fastapi.dev.router import router as dev_router from app.api.fastapi.exceptions import add_exception_handlers from app.api.fastapi.items.router import router as items_router from app.api.fastapi.lifespan import lifespan_factory @@ -11,5 +12,6 @@ def create_fastapi_app(settings: Settings) -> FastAPI: add_exception_handlers(app=app) app.include_router(items_router) + app.include_router(dev_router) return app diff --git a/app/api/fastapi/dev/__init__.py b/app/api/fastapi/dev/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/api/fastapi/dev/router.py b/app/api/fastapi/dev/router.py new file mode 100644 index 0000000..a423bf5 --- /dev/null +++ b/app/api/fastapi/dev/router.py @@ -0,0 +1,23 @@ +from typing import Annotated, Any, Literal + +from fastapi import APIRouter, Depends, status + +from app.api.fastapi.dependencies import get_context +from app.core.sqlalchemy.context import Context +from app.domain.dev.commands import ItemCreateError, create_item_error_command +from app.domain.items.entities import Item + +router = APIRouter(prefix="/dev") + + +@router.post("/error", response_model=Item, status_code=status.HTTP_201_CREATED) +def item_error( + item_create: ItemCreateError, + context: Annotated[Context, Depends(get_context)], + error_type: Literal["domain", "unexpected"] | None = None, +) -> Any: + return create_item_error_command( + context, + item_create=item_create, + error_type=error_type, + ) diff --git a/app/api/fastapi/exceptions.py b/app/api/fastapi/exceptions.py index 0a90ea9..a2b6d8e 100644 --- a/app/api/fastapi/exceptions.py +++ b/app/api/fastapi/exceptions.py @@ -1,6 +1,6 @@ from fastapi import FastAPI, status from starlette.requests import Request -from starlette.responses import JSONResponse, PlainTextResponse, Response +from starlette.responses import JSONResponse from app.api.utils import ERROR_MAPPING from app.domain.exceptions import ( @@ -9,16 +9,23 @@ def add_exception_handlers(app: FastAPI) -> None: + @app.exception_handler(Exception) + async def exception_handler(request: Request, exc: Exception) -> JSONResponse: + return JSONResponse( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + content={"detail": "Internal Server Error"}, + ) + @app.exception_handler(DomainError) - async def domain_exception_handler(request: Request, exc: DomainError) -> Response: + async def domain_exception_handler( + request: Request, + exc: DomainError, + ) -> JSONResponse: + status_code = status.HTTP_500_INTERNAL_SERVER_ERROR + for error_cls in type(exc).mro(): if issubclass(error_cls, DomainError) and error_cls in ERROR_MAPPING: - return JSONResponse( - status_code=ERROR_MAPPING[error_cls], - content={"detail": str(exc)}, - ) + status_code = ERROR_MAPPING[error_cls] + break - return PlainTextResponse( - status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - content="Internal Server Error", - ) + return JSONResponse(status_code=status_code, content={"detail": str(exc)}) diff --git a/app/api/flask/app.py b/app/api/flask/app.py index 3072af9..a01edcc 100644 --- a/app/api/flask/app.py +++ b/app/api/flask/app.py @@ -1,5 +1,6 @@ from flask import Flask +from app.api.flask.dev.router import router as dev_router from app.api.flask.exceptions import add_exception_handlers from app.api.flask.items.router import router as items_router from app.api.flask.utils import init_app @@ -11,6 +12,7 @@ def create_flask_app(settings: Settings) -> Flask: add_exception_handlers(app=app) app.register_blueprint(items_router) + app.register_blueprint(dev_router) init_app(settings=settings, app=app) return app diff --git a/app/api/flask/dependencies.py b/app/api/flask/dependencies.py index 3f39dd1..6221ec1 100644 --- a/app/api/flask/dependencies.py +++ b/app/api/flask/dependencies.py @@ -1,6 +1,7 @@ from flask import current_app, g from app.core.sqlalchemy.context import Context +from app.infrastructure.sqlalchemy.logger import logger def get_context() -> Context: @@ -12,10 +13,19 @@ def get_context() -> Context: def sql_session_teardown(error: BaseException | None) -> None: + # In Flask, `error` is only set for unhandled exceptions + # The global Exception handler catches everything, so `error` is always None. + sql_session = g.pop("sql_session", None) - if sql_session is not None: - if error is not None: - sql_session.rollback() - else: - sql_session.commit() + if not sql_session: + return + + if g.pop("error", None): + logger.error(f"Rollback due to '{g.exception}'") + sql_session.rollback() sql_session.close() + return + + sql_session.commit() + logger.info("Commit ok") + sql_session.close() diff --git a/app/api/flask/dev/__init__.py b/app/api/flask/dev/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/api/flask/dev/router.py b/app/api/flask/dev/router.py new file mode 100644 index 0000000..9aca7cf --- /dev/null +++ b/app/api/flask/dev/router.py @@ -0,0 +1,23 @@ +from flask import Blueprint, Response, request + +from app.api.flask.dependencies import get_context +from app.domain.dev.commands import ItemCreateError, create_item_error_command + +router = Blueprint("dev", __name__, url_prefix="/dev") + + +@router.post("/error") +def item_error() -> Response: + error_type = request.args["error_type"] + context = get_context() + item_create = ItemCreateError.model_validate(request.get_json()) + item = create_item_error_command( + context, + item_create=item_create, + error_type=error_type, # ty:ignore[invalid-argument-type] + ) + return Response( + response=item.model_dump_json(), + status=201, + content_type="application/json", + ) diff --git a/app/api/flask/exceptions.py b/app/api/flask/exceptions.py index f70f249..ac137ea 100644 --- a/app/api/flask/exceptions.py +++ b/app/api/flask/exceptions.py @@ -1,24 +1,55 @@ -from flask import Flask, Response +import json + +from flask import Flask, Response, g from pydantic import ValidationError +from werkzeug.exceptions import HTTPException from app.api.utils import ERROR_MAPPING from app.domain.exceptions import DomainError def add_exception_handlers(app: Flask) -> None: + @app.errorhandler(Exception) + def exception_handler(exc: Exception) -> Response: + # Ensure rollback in teardown + g.error = True + g.exception = exc + + return Response( + response=json.dumps({"detail": "Internal Server Error"}), + status=500, + content_type="application/json", + ) + + @app.errorhandler(HTTPException) + def http_exception_handler(exc: HTTPException) -> Response: + # Ensure rollback in teardown + g.error = True + g.exception = exc + + return Response( + response=json.dumps({"detail": exc.description}), + status=exc.code, + content_type="application/json", + ) + @app.errorhandler(DomainError) def domain_exception_handler(exc: DomainError) -> Response: + # Ensure rollback in teardown + g.error = True + g.exception = exc + + status_code = 500 + for error_cls in type(exc).mro(): if issubclass(error_cls, DomainError) and error_cls in ERROR_MAPPING: - return Response( - response="{'detail': str(exc)}", - status=ERROR_MAPPING[error_cls], - ) + status_code = ERROR_MAPPING[error_cls] + break return Response( - response="Internal Server Error", - status=500, - content_type="text/plain", + response=json.dumps({"detail": str(exc)}), + status=status_code, + content_type="application/json", ) @app.errorhandler(ValidationError) diff --git a/app/core/fastapi.py b/app/core/fastapi.py new file mode 100644 index 0000000..506c25b --- /dev/null +++ b/app/core/fastapi.py @@ -0,0 +1,5 @@ +from app.api.fastapi.app import create_fastapi_app +from app.core.settings import Settings + +settings = Settings() # ty:ignore[missing-argument] +app = create_fastapi_app(settings=settings) diff --git a/app/core/settings.py b/app/core/settings.py index f6b4bc9..e9c90f0 100644 --- a/app/core/settings.py +++ b/app/core/settings.py @@ -8,7 +8,7 @@ class DjangoSettings(BaseModel): - debug: bool = True + debug: bool = False secret_key: str root_urlconf: str = "app.api.django.urls" installed_apps: list[str] = DJANGO_APPS diff --git a/app/domain/dev/__init__.py b/app/domain/dev/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/domain/dev/commands.py b/app/domain/dev/commands.py new file mode 100644 index 0000000..c56bde9 --- /dev/null +++ b/app/domain/dev/commands.py @@ -0,0 +1,34 @@ +from typing import Literal + +from pydantic import BaseModel + +from app.domain.context import ContextProtocol +from app.domain.entities import EntityId +from app.domain.exceptions import BadRequestError +from app.domain.items.entities import Item + + +class UnexpectedError(Exception): + pass + + +class ItemCreateError(BaseModel): + id: EntityId + + +def create_item_error_command( + context: ContextProtocol, + /, + item_create: ItemCreateError, + error_type: Literal["domain", "unexpected"] | None = None, +) -> Item: + item = Item(id=item_create.id, name="Item", description="Wonderful item") + context.item_repository.save(item) + + if error_type == "domain": + raise BadRequestError("Bad Request") + + if error_type == "unexpected": + raise UnexpectedError() + + return item diff --git a/app/domain/dev/entities.py b/app/domain/dev/entities.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/api/clients/django.py b/tests/api/clients/django.py index f67cab5..1a2d3b2 100644 --- a/tests/api/clients/django.py +++ b/tests/api/clients/django.py @@ -16,6 +16,8 @@ def post(self, *args: Any, **kwargs: Any) -> Any: if "json" in kwargs: kwargs["data"] = kwargs.pop("json") kwargs["content_type"] = "application/json" + if "params" in kwargs: + kwargs["query_params"] = kwargs.pop("params") return self._client.post(*args, **kwargs) def patch(self, *args: Any, **kwargs: Any) -> Any: diff --git a/tests/api/clients/flask.py b/tests/api/clients/flask.py index 0d7f19f..3624dde 100644 --- a/tests/api/clients/flask.py +++ b/tests/api/clients/flask.py @@ -1,12 +1,13 @@ from typing import Any from flask.testing import FlaskClient +from werkzeug.test import TestResponse from tests.api.clients.base import HTTPClient class WrappedFlaskResponse: - def __init__(self, response: Any) -> None: + def __init__(self, response: TestResponse) -> None: self._response = response @property @@ -25,6 +26,8 @@ def get(self, *args: Any, **kwargs: Any) -> Any: return WrappedFlaskResponse(self._client.get(*args, **kwargs)) def post(self, *args: Any, **kwargs: Any) -> Any: + if "params" in kwargs: + kwargs["query_string"] = kwargs.pop("params") return WrappedFlaskResponse(self._client.post(*args, **kwargs)) def patch(self, *args: Any, **kwargs: Any) -> Any: diff --git a/tests/api/conftest.py b/tests/api/conftest.py index aa59bdf..1866a3c 100644 --- a/tests/api/conftest.py +++ b/tests/api/conftest.py @@ -53,7 +53,7 @@ def client( django_setup: None, ) -> Iterator[HTTPClient]: if request.param == "fastapi": - with TestClient(fastapi_app) as client: + with TestClient(fastapi_app, raise_server_exceptions=False) as client: yield client elif request.param == "flask": yield WrappedFlaskClient(flask_app.test_client()) diff --git a/tests/api/items/test_create_item.py b/tests/api/items/test_create_item.py index 794e35d..b0ccc49 100644 --- a/tests/api/items/test_create_item.py +++ b/tests/api/items/test_create_item.py @@ -1,12 +1,20 @@ +import uuid + import pytest +from sqlalchemy.orm import Session from starlette import status +from app.infrastructure.sqlalchemy.models.items import SQLItemModel from tests.api.clients.base import HTTPClient from tests.factories.items import ItemFactory @pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) -def test_create_item(item_factory: ItemFactory, client: HTTPClient) -> None: +def test_create_item( + item_factory: ItemFactory, + client: HTTPClient, + session: Session, +) -> None: item = item_factory.build() data = item.model_dump(exclude={"id"}) @@ -18,6 +26,13 @@ def test_create_item(item_factory: ItemFactory, client: HTTPClient) -> None: assert result["name"] == item.name assert result["description"] == item.description + item_id = uuid.UUID(result["id"]) + item_db = session.get(SQLItemModel, item_id) + assert item_db + assert item_db.id == item_id + assert item_db.name == item.name + assert item_db.description == item.description + @pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) def test_create_item_invalid_data(client: HTTPClient) -> None: diff --git a/tests/api/items/test_delete_item.py b/tests/api/items/test_delete_item.py index 22d79d0..125f4c8 100644 --- a/tests/api/items/test_delete_item.py +++ b/tests/api/items/test_delete_item.py @@ -1,20 +1,29 @@ import uuid import pytest +from sqlalchemy.orm import Session from starlette import status +from app.infrastructure.sqlalchemy.models.items import SQLItemModel from tests.api.clients.base import HTTPClient from tests.factories.items import ItemFactory @pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) -def test_delete_item(item_factory: ItemFactory, client: HTTPClient) -> None: +def test_delete_item( + item_factory: ItemFactory, + client: HTTPClient, + session: Session, +) -> None: item = item_factory.create_one() response = client.delete(f"/items/{item.id}") assert response.status_code == status.HTTP_204_NO_CONTENT + item_db = session.get(SQLItemModel, item.id) + assert item_db is None + @pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) def test_delete_item_not_found(client: HTTPClient) -> None: diff --git a/tests/api/items/test_items_special.py b/tests/api/items/test_items_special.py new file mode 100644 index 0000000..37ba13a --- /dev/null +++ b/tests/api/items/test_items_special.py @@ -0,0 +1,49 @@ +import uuid + +import pytest +from fastapi import status +from sqlalchemy.orm import Session + +from app.infrastructure.sqlalchemy.models.items import SQLItemModel +from tests.api.clients.base import HTTPClient + + +@pytest.mark.parametrize( + "error_type, status_code, error_message", + [ + ("domain", status.HTTP_400_BAD_REQUEST, "Bad Request"), + ("unexpected", status.HTTP_500_INTERNAL_SERVER_ERROR, "Internal Server Error"), + ], +) +@pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) +def test_item_error( + client: HTTPClient, + session: Session, + error_type: str, + status_code: int, + error_message: str, +) -> None: + item_id = uuid.uuid7() + + response = client.post( + "/dev/error", + params={"error_type": error_type}, + json={"id": str(item_id)}, + ) + + assert response.status_code == status_code + result = response.json() + assert result["detail"] == error_message + + # Check rollback after error + item_db = session.get(SQLItemModel, item_id) + assert item_db is None + + +@pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) +def test_http_error(client: HTTPClient) -> None: + response = client.get("/unknown") + + assert response.status_code == status.HTTP_404_NOT_FOUND + result = response.json() + assert "not found" in result["detail"].lower() diff --git a/tests/api/items/test_update_item.py b/tests/api/items/test_update_item.py index 2086889..4595a2f 100644 --- a/tests/api/items/test_update_item.py +++ b/tests/api/items/test_update_item.py @@ -1,14 +1,20 @@ import uuid import pytest +from sqlalchemy.orm import Session from starlette import status +from app.infrastructure.sqlalchemy.models.items import SQLItemModel from tests.api.clients.base import HTTPClient from tests.factories.items import ItemFactory @pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) -def test_update_item(item_factory: ItemFactory, client: HTTPClient) -> None: +def test_update_item( + item_factory: ItemFactory, + client: HTTPClient, + session: Session, +) -> None: updated_name = "New name" updated_description = "New description" data = {"name": updated_name, "description": updated_description} @@ -22,6 +28,12 @@ def test_update_item(item_factory: ItemFactory, client: HTTPClient) -> None: assert result["name"] == updated_name assert result["description"] == updated_description + item_db = session.get(SQLItemModel, item.id) + assert item_db + assert item_db.id == item.id + assert item_db.name == updated_name + assert item_db.description == updated_description + @pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) def test_update_item_name(item_factory: ItemFactory, client: HTTPClient) -> None: diff --git a/tests/logger.py b/tests/logger.py new file mode 100644 index 0000000..f09ea6b --- /dev/null +++ b/tests/logger.py @@ -0,0 +1,3 @@ +import logging + +logger = logging.getLogger("tests") diff --git a/tests/plugins/database.py b/tests/plugins/database.py index f0b38b2..543c24e 100644 --- a/tests/plugins/database.py +++ b/tests/plugins/database.py @@ -10,7 +10,10 @@ @pytest.fixture(scope="session") def engine(app_settings: Settings) -> Engine: - engine = create_engine(url=str(app_settings.postgres_dsn)) + engine = create_engine( + url=str(app_settings.postgres_dsn), + **app_settings.postgres_params, + ) OrmEntity.metadata.drop_all(engine) OrmEntity.metadata.create_all(engine)