From 6d32ae008636b35405e59fc914f79b8843a1b088 Mon Sep 17 00:00:00 2001 From: Julien Pillaud Date: Tue, 2 Jun 2026 09:08:57 +0200 Subject: [PATCH] feat: dependencies override --- app/api/dependencies.py | 8 ++++ app/api/django/dependencies.py | 10 ++++- app/api/django/items/routes.py | 11 ++++- app/api/django/urls.py | 10 ++++- app/api/fastapi/dependencies.py | 17 +++---- app/api/fastapi/dev/router.py | 8 ++++ app/api/fastapi/lifespan.py | 7 ++- app/api/flask/dependencies.py | 12 +++-- app/api/flask/dev/router.py | 11 +++++ app/api/flask/utils.py | 7 ++- app/core/django/context.py | 9 ++++ app/core/settings.py | 9 ++++ app/core/sqlalchemy/context.py | 9 +++- app/infrastructure/sqlalchemy/utils.py | 61 ++++++++++++++------------ tests/api/conftest.py | 5 ++- tests/api/items/test_items_special.py | 10 +++++ tests/conftest.py | 3 +- tests/plugins/database.py | 26 ++++------- 18 files changed, 156 insertions(+), 77 deletions(-) create mode 100644 app/api/dependencies.py diff --git a/app/api/dependencies.py b/app/api/dependencies.py new file mode 100644 index 0000000..67de102 --- /dev/null +++ b/app/api/dependencies.py @@ -0,0 +1,8 @@ +from functools import lru_cache + +from app.core.settings import Settings + + +@lru_cache +def get_settings() -> Settings: + return Settings() # ty: ignore[missing-argument] diff --git a/app/api/django/dependencies.py b/app/api/django/dependencies.py index ae4d7d7..a90585c 100644 --- a/app/api/django/dependencies.py +++ b/app/api/django/dependencies.py @@ -1,5 +1,11 @@ +from typing import Annotated + +from fast_depends import Depends + +from app.api.dependencies import get_settings from app.core.django.context import Context +from app.core.settings import Settings -def get_context() -> Context: - return Context() +def get_context(settings: Annotated[Settings, Depends(get_settings)]) -> Context: + return Context(settings=settings) diff --git a/app/api/django/items/routes.py b/app/api/django/items/routes.py index 5f94c52..2ee33f3 100644 --- a/app/api/django/items/routes.py +++ b/app/api/django/items/routes.py @@ -80,7 +80,16 @@ def delete( return JsonResponse(data="", status=204, safe=False) -class ItemViewSpecial(View): +class DevEnvView(View): + @inject + def get(self, context: Annotated[Context, Depends(get_context)]) -> JsonResponse: + return JsonResponse( + data={"environment": context.environment}, + safe=False, + ) + + +class DevErrorView(View): body_models: ClassVar[dict[str, type[BaseModel]]] = {"POST": ItemCreateError} @inject(cast=False) diff --git a/app/api/django/urls.py b/app/api/django/urls.py index 9feb1f8..9176d3d 100644 --- a/app/api/django/urls.py +++ b/app/api/django/urls.py @@ -2,12 +2,18 @@ from django.urls.resolvers import URLPattern, URLResolver from app.api.django.handlers import custom_handler404 -from app.api.django.items.routes import ItemView, ItemViewDetail, ItemViewSpecial +from app.api.django.items.routes import ( + DevEnvView, + DevErrorView, + ItemView, + ItemViewDetail, +) urlpatterns: list[URLPattern | URLResolver] = [ path("items", ItemView.as_view()), path("items/", ItemViewDetail.as_view()), - path("dev/error", ItemViewSpecial.as_view()), + path("dev/env", DevEnvView.as_view()), + path("dev/error", DevErrorView.as_view()), ] # Override handler to return JSON response diff --git a/app/api/fastapi/dependencies.py b/app/api/fastapi/dependencies.py index 42a6bdb..4301698 100644 --- a/app/api/fastapi/dependencies.py +++ b/app/api/fastapi/dependencies.py @@ -1,25 +1,22 @@ from collections.abc import Iterator -from functools import lru_cache from typing import Annotated from fastapi import Depends from sqlalchemy.orm import Session from starlette.requests import Request +from app.api.dependencies import get_settings from app.core.settings import Settings from app.core.sqlalchemy.context import Context -from app.infrastructure.sqlalchemy.utils import managed_session - - -@lru_cache -def get_settings() -> Settings: - return Settings() # ty: ignore[missing-argument] def get_sql_session(request: Request) -> Iterator[Session]: - with managed_session(request.app.state.sql_session_factory) as session: + with request.app.state.sql_resource.session() as session: yield session -def get_context(sql_session: Annotated[Session, Depends(get_sql_session)]) -> Context: - return Context(sql_session=sql_session) +def get_context( + settings: Annotated[Settings, Depends(get_settings)], + sql_session: Annotated[Session, Depends(get_sql_session)], +) -> Context: + return Context(settings=settings, sql_session=sql_session) diff --git a/app/api/fastapi/dev/router.py b/app/api/fastapi/dev/router.py index a423bf5..5580a68 100644 --- a/app/api/fastapi/dev/router.py +++ b/app/api/fastapi/dev/router.py @@ -3,6 +3,7 @@ from fastapi import APIRouter, Depends, status from app.api.fastapi.dependencies import get_context +from app.core.settings import AppEnvironment 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 @@ -10,6 +11,13 @@ router = APIRouter(prefix="/dev") +@router.get("/env") +def get_env( + context: Annotated[Context, Depends(get_context)], +) -> dict[str, AppEnvironment]: + return {"environment": context.environment} + + @router.post("/error", response_model=Item, status_code=status.HTTP_201_CREATED) def item_error( item_create: ItemCreateError, diff --git a/app/api/fastapi/lifespan.py b/app/api/fastapi/lifespan.py index 9497be9..9cdf4e2 100644 --- a/app/api/fastapi/lifespan.py +++ b/app/api/fastapi/lifespan.py @@ -5,7 +5,7 @@ from app.api.logger import logger from app.core.settings import Settings -from app.infrastructure.sqlalchemy.utils import create_sql_resource +from app.infrastructure.sqlalchemy.utils import SQLResource def lifespan_factory( @@ -14,9 +14,8 @@ def lifespan_factory( @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncIterator[None]: - sql_resource = create_sql_resource(settings=settings) - app.state.sql_engine = sql_resource.engine - app.state.sql_session_factory = sql_resource.session_factory + sql_resource = SQLResource.from_settings(settings) + app.state.sql_resource = sql_resource logger.info("Application startup complete") yield diff --git a/app/api/flask/dependencies.py b/app/api/flask/dependencies.py index 2723a86..bcabab0 100644 --- a/app/api/flask/dependencies.py +++ b/app/api/flask/dependencies.py @@ -5,14 +5,18 @@ from flask import current_app from sqlalchemy.orm import Session +from app.api.dependencies import get_settings +from app.core.settings import Settings from app.core.sqlalchemy.context import Context -from app.infrastructure.sqlalchemy.utils import managed_session def get_sql_session() -> Iterator[Session]: - with managed_session(current_app.config["SQL_SESSION_FACTORY"]) as session: + with current_app.config["SQL_RESOURCE"].session() as session: yield session -def get_context(sql_session: Annotated[Session, Depends(get_sql_session)]) -> Context: - return Context(sql_session=sql_session) +def get_context( + settings: Annotated[Settings, Depends(get_settings)], + sql_session: Annotated[Session, Depends(get_sql_session)], +) -> Context: + return Context(settings=settings, sql_session=sql_session) diff --git a/app/api/flask/dev/router.py b/app/api/flask/dev/router.py index e4547e6..d4cdc91 100644 --- a/app/api/flask/dev/router.py +++ b/app/api/flask/dev/router.py @@ -1,3 +1,4 @@ +import json from typing import Annotated from fast_depends import Depends, inject @@ -10,6 +11,16 @@ router = Blueprint("dev", __name__, url_prefix="/dev") +@router.get("/env") +@inject +def get_env(context: Annotated[Context, Depends(get_context)]) -> Response: + return Response( + json.dumps({"environment": context.environment}), + status=200, + content_type="application/json", + ) + + @router.post("/error") @inject def item_error( diff --git a/app/api/flask/utils.py b/app/api/flask/utils.py index 61691ea..fd6a9f2 100644 --- a/app/api/flask/utils.py +++ b/app/api/flask/utils.py @@ -1,10 +1,9 @@ from flask import Flask from app.core.settings import Settings -from app.infrastructure.sqlalchemy.utils import create_sql_resource +from app.infrastructure.sqlalchemy.utils import SQLResource 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 + sql_resource = SQLResource.from_settings(settings) + app.config["SQL_RESOURCE"] = sql_resource diff --git a/app/core/django/context.py b/app/core/django/context.py index 10a72c5..733dcb4 100644 --- a/app/core/django/context.py +++ b/app/core/django/context.py @@ -1,11 +1,20 @@ from functools import cached_property +from app.core.settings import AppEnvironment, Settings from app.domain.context import ContextProtocol from app.domain.items.repository import ItemRepositoryProtocol from app.infrastructure.django.items import ItemRepository class Context(ContextProtocol): + def __init__(self, settings: Settings) -> None: + self.settings = settings + + @property + def environment(self) -> AppEnvironment: + # property to test dependencies override + return self.settings.environment + @cached_property def item_repository(self) -> ItemRepositoryProtocol: return ItemRepository() diff --git a/app/core/settings.py b/app/core/settings.py index e9c90f0..3124384 100644 --- a/app/core/settings.py +++ b/app/core/settings.py @@ -1,3 +1,4 @@ +from enum import StrEnum from typing import Any from pydantic import BaseModel, PostgresDsn, SecretStr, computed_field @@ -7,6 +8,12 @@ from app.infrastructure.django.apps import DJANGO_APPS +class AppEnvironment(StrEnum): + DEVELOPMENT = "development" + TESTING = "testing" + PRODUCTION = "production" + + class DjangoSettings(BaseModel): debug: bool = False secret_key: str @@ -24,6 +31,8 @@ class Settings(BaseSettings): nested_model_default_partial_update=True, ) + environment: AppEnvironment + django: DjangoSettings postgres_user: str diff --git a/app/core/sqlalchemy/context.py b/app/core/sqlalchemy/context.py index c3e17a4..3218b4c 100644 --- a/app/core/sqlalchemy/context.py +++ b/app/core/sqlalchemy/context.py @@ -2,15 +2,22 @@ from sqlalchemy.orm import Session +from app.core.settings import AppEnvironment, Settings from app.domain.context import ContextProtocol from app.domain.items.repository import ItemRepositoryProtocol from app.infrastructure.sqlalchemy.items import SQLItemRepository class Context(ContextProtocol): - def __init__(self, sql_session: Session) -> None: + def __init__(self, settings: Settings, sql_session: Session) -> None: + self.settings = settings self.sql_session = sql_session + @property + def environment(self) -> AppEnvironment: + # property to test dependencies override + return self.settings.environment + @cached_property def item_repository(self) -> ItemRepositoryProtocol: return SQLItemRepository(session=self.sql_session) diff --git a/app/infrastructure/sqlalchemy/utils.py b/app/infrastructure/sqlalchemy/utils.py index 943711e..67b7882 100644 --- a/app/infrastructure/sqlalchemy/utils.py +++ b/app/infrastructure/sqlalchemy/utils.py @@ -16,6 +16,38 @@ class SQLResource(BaseModel): engine: Engine session_factory: sessionmaker[Session] + @classmethod + def from_settings(cls, settings: Settings, /) -> SQLResource: + engine = create_engine( + url=str(settings.postgres_dsn), + **settings.postgres_params, + ) + with engine.connect() as connection: + connection.execute(text("SELECT 1")) + logger.info("SQL engine up") + return cls( + engine=engine, + session_factory=sessionmaker(bind=engine), + ) + + def create_all(self) -> None: + OrmEntity.metadata.drop_all(self.engine) + OrmEntity.metadata.create_all(self.engine) + + @contextmanager + def session(self) -> Iterator[Session]: + _session = self.session_factory() + try: + yield _session + _session.commit() + logger.info("Commit ok") + except Exception as error: + logger.error(f"Rollback due to {error}") + _session.rollback() + raise + finally: + _session.close() + def release(self) -> None: logger.info("SQL engine released") self.engine.dispose() @@ -25,32 +57,3 @@ def reset(self) -> None: for table in reversed(OrmEntity.metadata.sorted_tables): session.execute(table.delete()) session.commit() - - -def create_sql_resource(settings: Settings) -> SQLResource: - engine = create_engine( - url=str(settings.postgres_dsn), - **settings.postgres_params, - ) - with engine.connect() as connection: - connection.execute(text("SELECT 1")) - logger.info("SQL engine up") - return SQLResource( - engine=engine, - session_factory=sessionmaker(bind=engine), - ) - - -@contextmanager -def managed_session(session_factory: sessionmaker[Session]) -> Iterator[Session]: - session = session_factory() - try: - yield session - session.commit() - logger.info("Commit ok") - except Exception as error: - logger.error(f"Rollback due to {error}") - session.rollback() - raise - finally: - session.close() diff --git a/tests/api/conftest.py b/tests/api/conftest.py index 1866a3c..c67298b 100644 --- a/tests/api/conftest.py +++ b/tests/api/conftest.py @@ -4,12 +4,13 @@ import pytest from django.conf import settings from django.test import Client +from fast_depends import dependency_provider from fastapi import FastAPI from fastapi.testclient import TestClient from flask import Flask +from app.api.dependencies import get_settings from app.api.fastapi.app import create_fastapi_app -from app.api.fastapi.dependencies import get_settings from app.api.flask.app import create_flask_app from app.core.settings import Settings from tests.api.clients.base import HTTPClient @@ -28,6 +29,7 @@ def fastapi_app(app_settings: Settings) -> FastAPI: @pytest.fixture(scope="session") def flask_app(app_settings: Settings) -> Flask: app = create_flask_app(settings=app_settings) + dependency_provider.override(get_settings, settings_override_func) return app @@ -43,6 +45,7 @@ def django_setup(app_settings: Settings) -> None: ALLOWED_HOSTS=["testserver"], ) django.setup() + dependency_provider.override(get_settings, settings_override_func) @pytest.fixture diff --git a/tests/api/items/test_items_special.py b/tests/api/items/test_items_special.py index 4c6eecb..63379a4 100644 --- a/tests/api/items/test_items_special.py +++ b/tests/api/items/test_items_special.py @@ -4,10 +4,20 @@ from fastapi import status from sqlalchemy.orm import Session +from app.core.settings import AppEnvironment from app.infrastructure.sqlalchemy.models.items import SQLItemModel from tests.api.clients.base import HTTPClient +@pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) +def test_get_env(client: HTTPClient) -> None: + response = client.get("/dev/env") + + assert response.status_code == status.HTTP_200_OK + result = response.json() + assert result["environment"] == AppEnvironment.TESTING + + @pytest.mark.parametrize("client", ["fastapi", "flask", "django"], indirect=True) def test_item_domain_error(client: HTTPClient, session: Session) -> None: item_id = uuid.uuid7() diff --git a/tests/conftest.py b/tests/conftest.py index 7ab9497..2e9bd1c 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,7 +3,7 @@ import pytest from pydantic import SecretStr -from app.core.settings import DjangoSettings, Settings +from app.core.settings import AppEnvironment, DjangoSettings, Settings pytest_plugins = [ "tests.plugins.database", @@ -14,6 +14,7 @@ @lru_cache def settings_override_func() -> Settings: return Settings( + environment=AppEnvironment.TESTING, postgres_user="user", postgres_password=SecretStr("password"), postgres_host="localhost", diff --git a/tests/plugins/database.py b/tests/plugins/database.py index 543c24e..6ca8aad 100644 --- a/tests/plugins/database.py +++ b/tests/plugins/database.py @@ -1,32 +1,22 @@ from collections.abc import Iterator import pytest -from sqlalchemy import Engine, create_engine from sqlalchemy.orm import Session from app.core.settings import Settings -from app.infrastructure.sqlalchemy.models.base import OrmEntity +from app.infrastructure.sqlalchemy.utils import SQLResource @pytest.fixture(scope="session") -def engine(app_settings: Settings) -> Engine: - engine = create_engine( - url=str(app_settings.postgres_dsn), - **app_settings.postgres_params, - ) - - OrmEntity.metadata.drop_all(engine) - OrmEntity.metadata.create_all(engine) - - return engine +def sql_resource(app_settings: Settings) -> SQLResource: + sql_resource = SQLResource.from_settings(app_settings) + sql_resource.create_all() + return sql_resource @pytest.fixture -def session(engine: Engine) -> Iterator[Session]: - with Session(engine) as session: +def session(sql_resource: SQLResource) -> Iterator[Session]: + with sql_resource.session() as session: yield session - with Session(engine) as session: - for table in reversed(OrmEntity.metadata.sorted_tables): - session.execute(table.delete()) - session.commit() + sql_resource.reset()