mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 10:38:32 +08:00
test: use sqlite3 session in test_wraps (#38749)
This commit is contained in:
parent
7909c0c0df
commit
831443b45c
@ -1,8 +1,12 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
import controllers.console.explore.wraps as wraps_module
|
||||
import models.model as model_module
|
||||
from controllers.console.explore.error import (
|
||||
AppAccessDeniedError,
|
||||
TrialAppLimitExceeded,
|
||||
@ -16,66 +20,107 @@ from controllers.console.explore.wraps import (
|
||||
trial_feature_enable,
|
||||
user_allowed_to_access_app,
|
||||
)
|
||||
from models import AccountTrialAppRecord, App, AppMode, InstalledApp, TrialApp
|
||||
|
||||
|
||||
def test_installed_app_required_not_found():
|
||||
@installed_app_required
|
||||
def view(installed_app):
|
||||
return "ok"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(), "tenant-1"),
|
||||
),
|
||||
patch("controllers.console.explore.wraps.db.session.scalar") as scalar_mock,
|
||||
):
|
||||
scalar_mock.return_value = None
|
||||
|
||||
with pytest.raises(NotFound):
|
||||
view("app-id")
|
||||
def _bind_database(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
monkeypatch.setattr(wraps_module.db, "session", sqlite_session)
|
||||
monkeypatch.setattr(model_module.db, "session", sqlite_session)
|
||||
|
||||
|
||||
def test_installed_app_required_app_deleted():
|
||||
installed_app = MagicMock(app=None)
|
||||
def _app() -> App:
|
||||
app = App(
|
||||
tenant_id=str(uuid4()),
|
||||
name="Explore App",
|
||||
mode=AppMode.CHAT,
|
||||
enable_site=True,
|
||||
enable_api=True,
|
||||
)
|
||||
app.id = str(uuid4())
|
||||
return app
|
||||
|
||||
|
||||
def _installed_app(*, app_id: str, tenant_id: str) -> InstalledApp:
|
||||
return InstalledApp(
|
||||
tenant_id=tenant_id,
|
||||
app_id=app_id,
|
||||
app_owner_tenant_id=str(uuid4()),
|
||||
position=0,
|
||||
is_pinned=False,
|
||||
last_used_at=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(InstalledApp, App)], indirect=True)
|
||||
def test_installed_app_required_not_found(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
tenant_id = str(uuid4())
|
||||
_bind_database(monkeypatch, sqlite_session)
|
||||
|
||||
@installed_app_required
|
||||
def view(installed_app):
|
||||
return "ok"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(), "tenant-1"),
|
||||
),
|
||||
patch("controllers.console.explore.wraps.db.session.scalar") as scalar_mock,
|
||||
patch("controllers.console.explore.wraps.db.session.delete"),
|
||||
patch("controllers.console.explore.wraps.db.session.commit"),
|
||||
with patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(), tenant_id),
|
||||
):
|
||||
scalar_mock.return_value = installed_app
|
||||
|
||||
with pytest.raises(NotFound):
|
||||
view("app-id")
|
||||
view(str(uuid4()))
|
||||
|
||||
|
||||
def test_installed_app_required_success():
|
||||
installed_app = MagicMock(app=MagicMock())
|
||||
@pytest.mark.parametrize("sqlite_session", [(InstalledApp, App)], indirect=True)
|
||||
def test_installed_app_required_app_deleted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
tenant_id = str(uuid4())
|
||||
installed_app = _installed_app(app_id=str(uuid4()), tenant_id=tenant_id)
|
||||
sqlite_session.add(installed_app)
|
||||
sqlite_session.commit()
|
||||
installed_app_id = installed_app.id
|
||||
_bind_database(monkeypatch, sqlite_session)
|
||||
|
||||
@installed_app_required
|
||||
def view(installed_app):
|
||||
return "ok"
|
||||
|
||||
with patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(), tenant_id),
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
view(installed_app_id)
|
||||
|
||||
assert sqlite_session.get(InstalledApp, installed_app_id) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(InstalledApp, App)], indirect=True)
|
||||
def test_installed_app_required_success(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
app = _app()
|
||||
installed_app = _installed_app(app_id=app.id, tenant_id=app.tenant_id)
|
||||
sqlite_session.add_all([app, installed_app])
|
||||
sqlite_session.commit()
|
||||
_bind_database(monkeypatch, sqlite_session)
|
||||
|
||||
@installed_app_required
|
||||
def view(installed_app):
|
||||
return installed_app
|
||||
|
||||
with (
|
||||
patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(), "tenant-1"),
|
||||
),
|
||||
patch("controllers.console.explore.wraps.db.session.scalar") as scalar_mock,
|
||||
with patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(), app.tenant_id),
|
||||
):
|
||||
scalar_mock.return_value = installed_app
|
||||
result = view(installed_app.id)
|
||||
|
||||
result = view("app-id")
|
||||
assert result == installed_app
|
||||
assert result.id == installed_app.id
|
||||
assert result.app is not None
|
||||
assert result.app.id == app.id
|
||||
|
||||
|
||||
def test_user_allowed_to_access_app_denied():
|
||||
@ -133,70 +178,74 @@ def test_user_allowed_to_access_app_success():
|
||||
assert view(installed_app) == "ok"
|
||||
|
||||
|
||||
def test_trial_app_required_not_allowed():
|
||||
@pytest.mark.parametrize("sqlite_session", [(TrialApp, App, AccountTrialAppRecord)], indirect=True)
|
||||
def test_trial_app_required_not_allowed(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
_bind_database(monkeypatch, sqlite_session)
|
||||
|
||||
@trial_app_required
|
||||
def view(app):
|
||||
return "ok"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(id="user-1"), None),
|
||||
),
|
||||
patch("controllers.console.explore.wraps.db.session.scalar") as scalar_mock,
|
||||
with patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(id=str(uuid4())), None),
|
||||
):
|
||||
scalar_mock.return_value = None
|
||||
|
||||
with pytest.raises(TrialAppNotAllowed):
|
||||
view("app-id")
|
||||
view(str(uuid4()))
|
||||
|
||||
|
||||
def test_trial_app_required_limit_exceeded():
|
||||
trial_app = MagicMock(trial_limit=1, app=MagicMock())
|
||||
record = MagicMock(count=1)
|
||||
@pytest.mark.parametrize("sqlite_session", [(TrialApp, App, AccountTrialAppRecord)], indirect=True)
|
||||
def test_trial_app_required_limit_exceeded(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
account_id = str(uuid4())
|
||||
app = _app()
|
||||
trial_app = TrialApp(app_id=app.id, tenant_id=app.tenant_id, trial_limit=1)
|
||||
record = AccountTrialAppRecord(account_id=account_id, app_id=app.id, count=1)
|
||||
sqlite_session.add_all([app, trial_app, record])
|
||||
sqlite_session.commit()
|
||||
_bind_database(monkeypatch, sqlite_session)
|
||||
|
||||
@trial_app_required
|
||||
def view(app):
|
||||
return "ok"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(id="user-1"), None),
|
||||
),
|
||||
patch("controllers.console.explore.wraps.db.session.scalar") as scalar_mock,
|
||||
with patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(id=account_id), None),
|
||||
):
|
||||
scalar_mock.side_effect = [
|
||||
trial_app,
|
||||
record,
|
||||
]
|
||||
|
||||
with pytest.raises(TrialAppLimitExceeded):
|
||||
view("app-id")
|
||||
view(app.id)
|
||||
|
||||
|
||||
def test_trial_app_required_success():
|
||||
trial_app = MagicMock(trial_limit=2, app=MagicMock())
|
||||
record = MagicMock(count=1)
|
||||
@pytest.mark.parametrize("sqlite_session", [(TrialApp, App, AccountTrialAppRecord)], indirect=True)
|
||||
def test_trial_app_required_success(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
account_id = str(uuid4())
|
||||
app = _app()
|
||||
trial_app = TrialApp(app_id=app.id, tenant_id=app.tenant_id, trial_limit=2)
|
||||
record = AccountTrialAppRecord(account_id=account_id, app_id=app.id, count=1)
|
||||
sqlite_session.add_all([app, trial_app, record])
|
||||
sqlite_session.commit()
|
||||
_bind_database(monkeypatch, sqlite_session)
|
||||
|
||||
@trial_app_required
|
||||
def view(app):
|
||||
return app
|
||||
|
||||
with (
|
||||
patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(id="user-1"), None),
|
||||
),
|
||||
patch("controllers.console.explore.wraps.db.session.scalar") as scalar_mock,
|
||||
with patch(
|
||||
"controllers.console.explore.wraps.current_account_with_tenant",
|
||||
return_value=(MagicMock(id=account_id), None),
|
||||
):
|
||||
scalar_mock.side_effect = [
|
||||
trial_app,
|
||||
record,
|
||||
]
|
||||
result = view(app.id)
|
||||
|
||||
result = view("app-id")
|
||||
assert result == trial_app.app
|
||||
assert result.id == app.id
|
||||
|
||||
|
||||
def test_trial_feature_enable_disabled():
|
||||
|
||||
Loading…
Reference in New Issue
Block a user