Skip to content
Merged
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
5 changes: 5 additions & 0 deletions app/api/django/dependencies.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
from app.core.django.context import Context


def get_context() -> Context:
return Context()
49 changes: 37 additions & 12 deletions app/api/django/items/routes.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -21,50 +23,73 @@
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)


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,
item_update=request.validated_data,
)
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)


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,
Expand Down
6 changes: 1 addition & 5 deletions app/api/django/middlewares.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
7 changes: 0 additions & 7 deletions app/api/fastapi/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
37 changes: 12 additions & 25 deletions app/api/flask/dependencies.py
Original file line number Diff line number Diff line change
@@ -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)
10 changes: 8 additions & 2 deletions app/api/flask/dev/router.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down
23 changes: 2 additions & 21 deletions app/api/flask/exceptions.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -9,24 +9,9 @@


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",
)

# Replace the default werkzeug HTML response
@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,
Expand All @@ -35,10 +20,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():
Expand Down
33 changes: 22 additions & 11 deletions app/api/flask/items/router.py
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -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("/<item_id>")
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(),
Expand All @@ -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(
Expand All @@ -47,8 +52,11 @@ def create_item() -> Response:


@router.patch("/<item_id>")
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(
Expand All @@ -59,8 +67,11 @@ def update_item(item_id: EntityId) -> Response:


@router.delete("/<item_id>")
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,
Expand Down
3 changes: 0 additions & 3 deletions app/api/flask/utils.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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)
5 changes: 5 additions & 0 deletions app/core/flask.py
Original file line number Diff line number Diff line change
@@ -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)
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
16 changes: 12 additions & 4 deletions tests/api/clients/flask.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading