From e6d598065ecdd14de540254f9e12decde98395e0 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Mon, 13 Jul 2026 15:20:40 +0900 Subject: [PATCH] test: use sqlite3 session in test_banner (#38746) --- .../console/explore/test_banner.py | 122 ++++++++++++------ 1 file changed, 82 insertions(+), 40 deletions(-) diff --git a/api/tests/unit_tests/controllers/console/explore/test_banner.py b/api/tests/unit_tests/controllers/console/explore/test_banner.py index 552dc7d8217..36b83151b72 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_banner.py +++ b/api/tests/unit_tests/controllers/console/explore/test_banner.py @@ -1,35 +1,76 @@ +from collections.abc import Iterator from datetime import datetime from inspect import unwrap -from unittest.mock import MagicMock, patch +from uuid import uuid4 +import pytest from flask import Flask +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session import controllers.console.explore.banner as banner_module +from models.base import TypeBase from models.enums import BannerStatus +from models.model import ExporleBanner + + +@pytest.fixture +def banner_session(sqlite_engine: Engine) -> Iterator[Session]: + """Create the banner table without its PostgreSQL-cast server defaults.""" + table = TypeBase.metadata.tables[ExporleBanner.__tablename__] + status_default = table.c.status.server_default + language_default = table.c.language.server_default + table.c.status.server_default = None + table.c.language.server_default = None + try: + TypeBase.metadata.create_all(sqlite_engine, tables=[table]) + finally: + table.c.status.server_default = status_default + table.c.language.server_default = language_default + + with Session(sqlite_engine, expire_on_commit=False) as session: + yield session + + +def _banner(*, text: str, language: str, link: str, created_at: datetime) -> ExporleBanner: + banner = ExporleBanner( + content={"text": text}, + link=link, + sort=1, + status=BannerStatus.ENABLED, + language=language, + ) + banner.id = str(uuid4()) + banner.created_at = created_at + return banner class TestBannerApi: - def test_get_banners_with_requested_language(self, app: Flask): + def test_get_banners_with_requested_language( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + banner_session: Session, + ): api = banner_module.BannerApi() method = unwrap(api.get) - banner = MagicMock() - banner.id = "b1" - banner.content = {"text": "hello"} - banner.link = "https://example.com" - banner.sort = 1 - banner.status = BannerStatus.ENABLED - banner.created_at = datetime(2024, 1, 1) + banner = _banner( + text="hello", + language="fr-FR", + link="https://example.com", + created_at=datetime(2024, 1, 1), + ) + banner_session.add(banner) + banner_session.commit() + monkeypatch.setattr(banner_module.db, "session", banner_session) - session = MagicMock() - session.scalars.return_value.all.return_value = [banner] - - with app.test_request_context("/?language=fr-FR"), patch.object(banner_module.db, "session", session): + with app.test_request_context("/?language=fr-FR"): result = method(api) assert result == [ { - "id": "b1", + "id": banner.id, "content": {"text": "hello"}, "link": "https://example.com", "sort": 1, @@ -38,49 +79,50 @@ class TestBannerApi: } ] - def test_get_banners_fallback_to_en_us(self, app: Flask): + def test_get_banners_fallback_to_en_us( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + banner_session: Session, + ): api = banner_module.BannerApi() method = unwrap(api.get) - banner = MagicMock() - banner.id = "b2" - banner.content = {"text": "fallback"} - banner.link = None - banner.sort = 1 - banner.status = BannerStatus.ENABLED - banner.created_at = None + banner = _banner( + text="fallback", + language="en-US", + link="https://example.com/fallback", + created_at=datetime(2024, 1, 2), + ) + banner_session.add(banner) + banner_session.commit() + monkeypatch.setattr(banner_module.db, "session", banner_session) - scalars_result = MagicMock() - scalars_result.all.side_effect = [ - [], - [banner], - ] - - session = MagicMock() - session.scalars.return_value = scalars_result - - with app.test_request_context("/?language=es-ES"), patch.object(banner_module.db, "session", session): + with app.test_request_context("/?language=es-ES"): result = method(api) assert result == [ { - "id": "b2", + "id": banner.id, "content": {"text": "fallback"}, - "link": None, + "link": "https://example.com/fallback", "sort": 1, "status": "enabled", - "created_at": None, + "created_at": "2024-01-02T00:00:00", } ] - def test_get_banners_default_language_en_us(self, app: Flask): + def test_get_banners_default_language_en_us( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + banner_session: Session, + ): api = banner_module.BannerApi() method = unwrap(api.get) + monkeypatch.setattr(banner_module.db, "session", banner_session) - session = MagicMock() - session.scalars.return_value.all.return_value = [] - - with app.test_request_context("/"), patch.object(banner_module.db, "session", session): + with app.test_request_context("/"): result = method(api) assert result == []