test: use sqlite3 session in test_banner (#38746)

This commit is contained in:
Asuka Minato 2026-07-13 15:20:40 +09:00 committed by GitHub
parent 9740c35f6f
commit e6d598065e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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 == []