mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 10:38:32 +08:00
test: use sqlite3 session in test_banner (#38746)
This commit is contained in:
parent
9740c35f6f
commit
e6d598065e
@ -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 == []
|
||||
|
||||
Loading…
Reference in New Issue
Block a user