From 0d45c2c7e5f564ff5c1f563d7a8f3fcf0afe9845 Mon Sep 17 00:00:00 2001 From: Julien Pillaud Date: Mon, 1 Jun 2026 17:38:21 +0200 Subject: [PATCH 1/2] feat: add fast-depends --- app/api/django/dependencies.py | 5 +++ app/api/django/items/routes.py | 49 ++++++++++++++++++++------- app/api/django/middlewares.py | 6 +--- app/api/fastapi/exceptions.py | 7 ---- app/api/flask/dependencies.py | 37 +++++++------------- app/api/flask/dev/router.py | 10 ++++-- app/api/flask/exceptions.py | 20 ++--------- app/api/flask/items/router.py | 33 ++++++++++++------ app/api/flask/utils.py | 3 -- app/core/flask.py | 5 +++ pyproject.toml | 1 + tests/api/clients/flask.py | 16 ++++++--- tests/api/items/test_items_special.py | 39 +++++++++++---------- uv.lock | 15 ++++++++ 14 files changed, 142 insertions(+), 104 deletions(-) create mode 100644 app/api/django/dependencies.py create mode 100644 app/core/flask.py diff --git a/app/api/django/dependencies.py b/app/api/django/dependencies.py new file mode 100644 index 0000000..ae4d7d7 --- /dev/null +++ b/app/api/django/dependencies.py @@ -0,0 +1,5 @@ +from app.core.django.context import Context + + +def get_context() -> Context: + return Context() diff --git a/app/api/django/items/routes.py b/app/api/django/items/routes.py index 7e99593..5f94c52 100644 --- a/app/api/django/items/routes.py +++ b/app/api/django/items/routes.py @@ -1,10 +1,12 @@ import uuid -from typing import ClassVar +from typing import Annotated, ClassVar from django.http import HttpRequest, JsonResponse from django.views import View +from fast_depends import Depends, inject from pydantic import BaseModel +from app.api.django.dependencies import get_context 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 @@ -21,13 +23,21 @@ class ItemView(View): body_models: ClassVar[dict[str, type[BaseModel]]] = {"POST": ItemCreate} - def get(self, request: HttpRequest) -> JsonResponse: - context = Context() + @inject + def get( + self, + request: HttpRequest, + context: Annotated[Context, Depends(get_context)], + ) -> JsonResponse: items = get_items_command(context) return JsonResponse(data=[item.model_dump() for item in items], safe=False) - def post(self, request: EnhancedHttpRequest[ItemCreate]) -> JsonResponse: - context = Context() + @inject(cast=False) + def post( + self, + request: EnhancedHttpRequest[ItemCreate], + context: Annotated[Context, Depends(get_context)], + ) -> JsonResponse: item = create_item_command(context, item_create=request.validated_data) return JsonResponse(item.model_dump(), status=201, safe=False) @@ -35,17 +45,23 @@ def post(self, request: EnhancedHttpRequest[ItemCreate]) -> JsonResponse: class ItemViewDetail(View): body_models: ClassVar[dict[str, type[BaseModel]]] = {"PATCH": ItemUpdate} - def get(self, request: HttpRequest, item_id: uuid.UUID) -> JsonResponse: - context = Context() + @inject + def get( + self, + request: HttpRequest, + item_id: uuid.UUID, + context: Annotated[Context, Depends(get_context)], + ) -> JsonResponse: item = get_item_command(context, item_id=item_id) return JsonResponse(data=item.model_dump(), safe=False) + @inject(cast=False) def patch( self, request: EnhancedHttpRequest[ItemUpdate], item_id: uuid.UUID, + context: Annotated[Context, Depends(get_context)], ) -> JsonResponse: - context = Context() item = update_item_command( context, item_id=item_id, @@ -53,8 +69,13 @@ def patch( ) return JsonResponse(data=item.model_dump(), safe=False) - def delete(self, request: HttpRequest, item_id: uuid.UUID) -> JsonResponse: - context = Context() + @inject + def delete( + self, + request: HttpRequest, + item_id: uuid.UUID, + context: Annotated[Context, Depends(get_context)], + ) -> JsonResponse: delete_item_command(context, item_id=item_id) return JsonResponse(data="", status=204, safe=False) @@ -62,9 +83,13 @@ def delete(self, request: HttpRequest, item_id: uuid.UUID) -> JsonResponse: class ItemViewSpecial(View): body_models: ClassVar[dict[str, type[BaseModel]]] = {"POST": ItemCreateError} - def post(self, request: EnhancedHttpRequest[ItemCreateError]) -> JsonResponse: + @inject(cast=False) + def post( + self, + request: EnhancedHttpRequest[ItemCreateError], + context: Annotated[Context, Depends(get_context)], + ) -> JsonResponse: error_type = request.GET.get("error_type") - context = Context() item = create_item_error_command( context, item_create=request.validated_data, diff --git a/app/api/django/middlewares.py b/app/api/django/middlewares.py index 126e76a..ed4d8e4 100644 --- a/app/api/django/middlewares.py +++ b/app/api/django/middlewares.py @@ -83,11 +83,7 @@ def process_exception( if isinstance(exc, DomainError): return handle_domain_exceptions(exc) - return HttpResponse( - content=json.dumps({"detail": "Internal Server Error"}), - status=500, - content_type="application/json", - ) + return HttpResponse(content="Internal Server Error", status=500) def handle_domain_exceptions(exc: DomainError) -> HttpResponse: diff --git a/app/api/fastapi/exceptions.py b/app/api/fastapi/exceptions.py index a2b6d8e..7dbd697 100644 --- a/app/api/fastapi/exceptions.py +++ b/app/api/fastapi/exceptions.py @@ -9,13 +9,6 @@ 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, diff --git a/app/api/flask/dependencies.py b/app/api/flask/dependencies.py index 6221ec1..2723a86 100644 --- a/app/api/flask/dependencies.py +++ b/app/api/flask/dependencies.py @@ -1,31 +1,18 @@ -from flask import current_app, g +from collections.abc import Iterator +from typing import Annotated -from app.core.sqlalchemy.context import Context -from app.infrastructure.sqlalchemy.logger import logger - - -def get_context() -> Context: - if "sql_session" not in g: - sql_session_factory = current_app.config["SQL_SESSION_FACTORY"] - g.sql_session = sql_session_factory() - - return Context(sql_session=g.sql_session) +from fast_depends import Depends +from flask import current_app +from sqlalchemy.orm import Session +from app.core.sqlalchemy.context import Context +from app.infrastructure.sqlalchemy.utils import managed_session -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 not sql_session: - return +def get_sql_session() -> Iterator[Session]: + with managed_session(current_app.config["SQL_SESSION_FACTORY"]) as session: + yield session - 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() +def get_context(sql_session: Annotated[Session, Depends(get_sql_session)]) -> Context: + return Context(sql_session=sql_session) diff --git a/app/api/flask/dev/router.py b/app/api/flask/dev/router.py index 9aca7cf..e4547e6 100644 --- a/app/api/flask/dev/router.py +++ b/app/api/flask/dev/router.py @@ -1,15 +1,21 @@ +from typing import Annotated + +from fast_depends import Depends, inject from flask import Blueprint, Response, request from app.api.flask.dependencies import get_context +from app.core.sqlalchemy.context import 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: +@inject +def item_error( + context: Annotated[Context, Depends(get_context)], +) -> Response: error_type = request.args["error_type"] - context = get_context() item_create = ItemCreateError.model_validate(request.get_json()) item = create_item_error_command( context, diff --git a/app/api/flask/exceptions.py b/app/api/flask/exceptions.py index ac137ea..8a92d49 100644 --- a/app/api/flask/exceptions.py +++ b/app/api/flask/exceptions.py @@ -1,6 +1,6 @@ import json -from flask import Flask, Response, g +from flask import Flask, Response from pydantic import ValidationError from werkzeug.exceptions import HTTPException @@ -11,22 +11,10 @@ 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", - ) + return Response(response="Internal Server Error", status=500) @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, @@ -35,10 +23,6 @@ def http_exception_handler(exc: HTTPException) -> Response: @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(): diff --git a/app/api/flask/items/router.py b/app/api/flask/items/router.py index 375cd2c..e8d0810 100644 --- a/app/api/flask/items/router.py +++ b/app/api/flask/items/router.py @@ -1,8 +1,10 @@ -from typing import Any +from typing import Annotated, Any +from fast_depends import Depends, inject from flask import Blueprint, Response, request from app.api.flask.dependencies import get_context +from app.core.sqlalchemy.context import Context from app.domain.entities import EntityId from app.domain.items.commands import ( create_item_command, @@ -17,15 +19,18 @@ @router.get("") -def get_items() -> Any: - context = get_context() +@inject +def get_items(context: Annotated[Context, Depends(get_context)]) -> Any: items = get_items_command(context) return [item.model_dump() for item in items] @router.get("/") -def get_item(item_id: EntityId) -> Response: - context = get_context() +@inject +def get_item( + item_id: EntityId, + context: Annotated[Context, Depends(get_context)], +) -> Response: item = get_item_command(context, item_id=item_id) return Response( response=item.model_dump_json(), @@ -35,8 +40,8 @@ def get_item(item_id: EntityId) -> Response: @router.post("") -def create_item() -> Response: - context = get_context() +@inject +def create_item(context: Annotated[Context, Depends(get_context)]) -> Response: item_create = ItemCreate.model_validate(request.get_json()) item = create_item_command(context, item_create=item_create) return Response( @@ -47,8 +52,11 @@ def create_item() -> Response: @router.patch("/") -def update_item(item_id: EntityId) -> Response: - context = get_context() +@inject +def update_item( + item_id: EntityId, + context: Annotated[Context, Depends(get_context)], +) -> Response: item_update = ItemUpdate.model_validate(request.get_json()) item = update_item_command(context, item_id=item_id, item_update=item_update) return Response( @@ -59,8 +67,11 @@ def update_item(item_id: EntityId) -> Response: @router.delete("/") -def delete_item(item_id: EntityId) -> Response: - context = get_context() +@inject +def delete_item( + item_id: EntityId, + context: Annotated[Context, Depends(get_context)], +) -> Response: delete_item_command(context, item_id=item_id) return Response( status=204, diff --git a/app/api/flask/utils.py b/app/api/flask/utils.py index 30d5ba3..61691ea 100644 --- a/app/api/flask/utils.py +++ b/app/api/flask/utils.py @@ -1,6 +1,5 @@ from flask import Flask -from app.api.flask.dependencies import sql_session_teardown from app.core.settings import Settings from app.infrastructure.sqlalchemy.utils import create_sql_resource @@ -9,5 +8,3 @@ def init_app(settings: Settings, app: Flask) -> None: sql_resource = create_sql_resource(settings=settings) app.config["SQL_ENGINE"] = sql_resource.engine app.config["SQL_SESSION_FACTORY"] = sql_resource.session_factory - - app.teardown_appcontext(sql_session_teardown) diff --git a/app/core/flask.py b/app/core/flask.py new file mode 100644 index 0000000..f2eeb7e --- /dev/null +++ b/app/core/flask.py @@ -0,0 +1,5 @@ +from app.api.flask.app import create_flask_app +from app.core.settings import Settings + +settings = Settings() # ty:ignore[missing-argument] +app = create_flask_app(settings=settings) diff --git a/pyproject.toml b/pyproject.toml index b9077fb..c5a2f8a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,6 +4,7 @@ version = "0.1.0" requires-python = ">=3.14" dependencies = [ "django==6.0.5", + "fast-depends==3.0.8", "fastapi[standard]==0.136.3", "flask==3.1.3", "gunicorn==26.0.0", diff --git a/tests/api/clients/flask.py b/tests/api/clients/flask.py index 3624dde..be57bbd 100644 --- a/tests/api/clients/flask.py +++ b/tests/api/clients/flask.py @@ -17,21 +17,29 @@ def status_code(self) -> int: def json(self) -> Any: return self._response.json + @property + def text(self) -> Any: + return self._response.text + class WrappedFlaskClient(HTTPClient): def __init__(self, client: FlaskClient) -> None: self._client = client def get(self, *args: Any, **kwargs: Any) -> Any: - return WrappedFlaskResponse(self._client.get(*args, **kwargs)) + response = self._client.get(*args, **kwargs) + return WrappedFlaskResponse(response=response) 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)) + response = self._client.post(*args, **kwargs) + return WrappedFlaskResponse(response=response) def patch(self, *args: Any, **kwargs: Any) -> Any: - return WrappedFlaskResponse(self._client.patch(*args, **kwargs)) + response = self._client.patch(*args, **kwargs) + return WrappedFlaskResponse(response=response) def delete(self, *args: Any, **kwargs: Any) -> Any: - return WrappedFlaskResponse(self._client.delete(*args, **kwargs)) + response = self._client.delete(*args, **kwargs) + return WrappedFlaskResponse(response=response) diff --git a/tests/api/items/test_items_special.py b/tests/api/items/test_items_special.py index 37ba13a..8b714ee 100644 --- a/tests/api/items/test_items_special.py +++ b/tests/api/items/test_items_special.py @@ -8,32 +8,37 @@ 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: +def test_item_domain_error(client: HTTPClient, session: Session) -> None: item_id = uuid.uuid7() response = client.post( "/dev/error", - params={"error_type": error_type}, + params={"error_type": "domain"}, json={"id": str(item_id)}, ) - assert response.status_code == status_code + assert response.status_code == status.HTTP_400_BAD_REQUEST result = response.json() - assert result["detail"] == error_message + assert result["detail"] == "Bad Request" + + # 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_item_unexpected_error(client: HTTPClient, session: Session) -> None: + item_id = uuid.uuid7() + + response = client.post( + "/dev/error", + params={"error_type": "unexpected"}, + json={"id": str(item_id)}, + ) + + assert response.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR + assert response.text == "Internal Server Error" # Check rollback after error item_db = session.get(SQLItemModel, item_id) diff --git a/uv.lock b/uv.lock index 020d2f5..7dc4d7a 100644 --- a/uv.lock +++ b/uv.lock @@ -186,6 +186,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/de/15/545e2b6cf2e3be84bc1ed85613edd75b8aea69807a71c26f4ca6a9258e82/email_validator-2.3.0-py3-none-any.whl", hash = "sha256:80f13f623413e6b197ae73bb10bf4eb0908faf509ad8362c5edeb0be7fd450b4", size = 35604, upload-time = "2025-08-26T13:09:05.858Z" }, ] +[[package]] +name = "fast-depends" +version = "3.0.8" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/f8/6d/787a21ca8043a8fdb737cf28f645e94a46fc30b44a31de54573299156bad/fast_depends-3.0.8.tar.gz", hash = "sha256:896b16f79a512b6ea1df721b0aa1708a192a06f964be6597e01fcf5412559101", size = 18382, upload-time = "2026-03-02T19:54:28.649Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/1d/1d/e4843e4eeb65f51447b8c22d200d12d8f94f27c97e77bb7162515cc8d61f/fast_depends-3.0.8-py3-none-any.whl", hash = "sha256:4c52c8a3907bca46d43e70e4364d6d016872d9a3aae4bc0c1c85e72e0a6a21c7", size = 25507, upload-time = "2026-03-02T19:54:27.594Z" }, +] + [[package]] name = "fastapi" version = "0.136.3" @@ -760,6 +773,7 @@ version = "0.1.0" source = { virtual = "." } dependencies = [ { name = "django" }, + { name = "fast-depends" }, { name = "fastapi", extra = ["standard"] }, { name = "flask" }, { name = "gunicorn" }, @@ -780,6 +794,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "django", specifier = "==6.0.5" }, + { name = "fast-depends", specifier = "==3.0.8" }, { name = "fastapi", extras = ["standard"], specifier = "==0.136.3" }, { name = "flask", specifier = "==3.1.3" }, { name = "gunicorn", specifier = "==26.0.0" }, From 9eb1c0d614b07b97ba2ca2bc69892d15857a57d2 Mon Sep 17 00:00:00 2001 From: Julien Pillaud Date: Mon, 1 Jun 2026 17:53:32 +0200 Subject: [PATCH 2/2] feat: add fast-depends --- app/api/flask/exceptions.py | 5 +---- tests/api/items/test_items_special.py | 21 ++++++++++++++++++++- 2 files changed, 21 insertions(+), 5 deletions(-) diff --git a/app/api/flask/exceptions.py b/app/api/flask/exceptions.py index 8a92d49..9d7d856 100644 --- a/app/api/flask/exceptions.py +++ b/app/api/flask/exceptions.py @@ -9,10 +9,7 @@ def add_exception_handlers(app: Flask) -> None: - @app.errorhandler(Exception) - def exception_handler(exc: Exception) -> Response: - return Response(response="Internal Server Error", status=500) - + # Replace the default werkzeug HTML response @app.errorhandler(HTTPException) def http_exception_handler(exc: HTTPException) -> Response: return Response( diff --git a/tests/api/items/test_items_special.py b/tests/api/items/test_items_special.py index 8b714ee..4c6eecb 100644 --- a/tests/api/items/test_items_special.py +++ b/tests/api/items/test_items_special.py @@ -27,7 +27,7 @@ def test_item_domain_error(client: HTTPClient, session: Session) -> None: assert item_db is None -@pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) +@pytest.mark.parametrize("client", ["fastapi", "django"], indirect=True) def test_item_unexpected_error(client: HTTPClient, session: Session) -> None: item_id = uuid.uuid7() @@ -45,6 +45,25 @@ def test_item_unexpected_error(client: HTTPClient, session: Session) -> None: assert item_db is None +@pytest.mark.parametrize("client", ["flask"], indirect=True) +def test_item_unexpected_error_flask(client: HTTPClient, session: Session) -> None: + item_id = uuid.uuid7() + + response = client.post( + "/dev/error", + params={"error_type": "unexpected"}, + json={"id": str(item_id)}, + ) + + assert response.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR + result = response.json() + assert "internal error" in result["detail"].lower() + + # 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")