mirror of
https://github.com/langgenius/dify.git
synced 2026-08-02 18:56:34 +08:00
test: use sqlite3 session in test_trial (#38732)
This commit is contained in:
parent
814fee5550
commit
b852252076
@ -9,6 +9,8 @@ from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import Forbidden, InternalServerError, NotFound
|
||||
|
||||
import controllers.console.explore.trial as module
|
||||
@ -36,7 +38,7 @@ from graphon.model_runtime.errors.invoke import InvokeError
|
||||
from graphon.variables import StringVariable
|
||||
from models import Account
|
||||
from models.account import TenantStatus
|
||||
from models.model import AppMode
|
||||
from models.model import AppMode, Site
|
||||
from services.app_ref_service import AppRef, MessageRef
|
||||
from services.errors.audio import SpeechToTextDisabledServiceError
|
||||
from services.errors.conversation import ConversationNotExistsError
|
||||
@ -45,6 +47,16 @@ from services.errors.llm import InvokeRateLimitError
|
||||
unwrap: Any = inspect_unwrap
|
||||
|
||||
|
||||
class _UsesSQLiteSession:
|
||||
sqlite_session: Session
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _provide_sqlite_session(self, sqlite_engine: Engine):
|
||||
with Session(sqlite_engine, expire_on_commit=False) as session:
|
||||
self.sqlite_session = session
|
||||
yield
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def account() -> Account:
|
||||
acc = Account(name="User", email="user@example.com")
|
||||
@ -58,6 +70,18 @@ def _file_data() -> Any:
|
||||
return file_data
|
||||
|
||||
|
||||
def _persist_site(sqlite_session: Session, app_id: str) -> Site:
|
||||
site = Site(
|
||||
app_id=app_id,
|
||||
title="Trial Site",
|
||||
default_language="en-US",
|
||||
customize_token_strategy="uuid",
|
||||
)
|
||||
sqlite_session.add(site)
|
||||
sqlite_session.commit()
|
||||
return site
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def trial_app_chat() -> MagicMock:
|
||||
app = MagicMock()
|
||||
@ -227,14 +251,14 @@ class TestTrialAppRemoteFileUploadApi:
|
||||
upload.assert_called_once_with(current_user=account, resource_tenant_id="app-tenant-id")
|
||||
|
||||
|
||||
class TestTrialAppWorkflowRunApi:
|
||||
class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
|
||||
def test_not_workflow_app(self, app: Flask, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
method = unwrap(api.post)
|
||||
|
||||
with app.test_request_context("/"):
|
||||
with pytest.raises(NotWorkflowAppError):
|
||||
method(api, MagicMock(), account, MagicMock(mode=AppMode.CHAT))
|
||||
method(api, self.sqlite_session, account, MagicMock(mode=AppMode.CHAT))
|
||||
|
||||
def test_success(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
@ -245,7 +269,7 @@ class TestTrialAppWorkflowRunApi:
|
||||
patch.object(module.AppGenerateService, "generate", return_value=MagicMock()),
|
||||
patch.object(module.RecommendedAppService, "add_trial_app_record"),
|
||||
):
|
||||
result = method(api, MagicMock(), account, trial_app_workflow)
|
||||
result = method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
assert result is not None
|
||||
|
||||
@ -262,7 +286,7 @@ class TestTrialAppWorkflowRunApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderNotInitializeError):
|
||||
method(api, MagicMock(), account, trial_app_workflow)
|
||||
method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
def test_workflow_quota_exceeded(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
@ -277,7 +301,7 @@ class TestTrialAppWorkflowRunApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderQuotaExceededError):
|
||||
method(api, MagicMock(), account, trial_app_workflow)
|
||||
method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
def test_workflow_model_not_support(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
@ -292,7 +316,7 @@ class TestTrialAppWorkflowRunApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderModelCurrentlyNotSupportError):
|
||||
method(api, MagicMock(), account, trial_app_workflow)
|
||||
method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
def test_workflow_invoke_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
@ -307,7 +331,7 @@ class TestTrialAppWorkflowRunApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(CompletionRequestError):
|
||||
method(api, MagicMock(), account, trial_app_workflow)
|
||||
method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
def test_workflow_rate_limit_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
@ -322,7 +346,7 @@ class TestTrialAppWorkflowRunApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(InvokeRateLimitHttpError):
|
||||
method(api, MagicMock(), account, trial_app_workflow)
|
||||
method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
def test_workflow_value_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
@ -337,7 +361,7 @@ class TestTrialAppWorkflowRunApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ValueError):
|
||||
method(api, MagicMock(), account, trial_app_workflow)
|
||||
method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
def test_workflow_generic_exception(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
|
||||
api = module.TrialAppWorkflowRunApi()
|
||||
@ -352,17 +376,17 @@ class TestTrialAppWorkflowRunApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(InternalServerError):
|
||||
method(api, MagicMock(), account, trial_app_workflow)
|
||||
method(api, self.sqlite_session, account, trial_app_workflow)
|
||||
|
||||
|
||||
class TestTrialChatApi:
|
||||
class TestTrialChatApi(_UsesSQLiteSession):
|
||||
def test_not_chat_app(self, app: Flask, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
method = unwrap(api.post)
|
||||
|
||||
with app.test_request_context("/", json={"inputs": {}, "query": "hi"}):
|
||||
with pytest.raises(NotChatAppError):
|
||||
method(api, MagicMock(), account, MagicMock(mode="completion"))
|
||||
method(api, self.sqlite_session, account, MagicMock(mode="completion"))
|
||||
|
||||
def test_success(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -373,7 +397,7 @@ class TestTrialChatApi:
|
||||
patch.object(module.AppGenerateService, "generate", return_value=MagicMock()),
|
||||
patch.object(module.RecommendedAppService, "add_trial_app_record"),
|
||||
):
|
||||
result = method(api, MagicMock(), account, trial_app_chat)
|
||||
result = method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
assert result is not None
|
||||
|
||||
@ -390,7 +414,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_conversation_completed(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -405,7 +429,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ConversationCompletedError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -420,7 +444,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(AppUnavailableError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -435,7 +459,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderNotInitializeError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -450,7 +474,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderQuotaExceededError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_model_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -465,7 +489,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderModelCurrentlyNotSupportError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -480,7 +504,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(CompletionRequestError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_rate_limit_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -495,7 +519,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(InvokeRateLimitHttpError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_value_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -510,7 +534,7 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ValueError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
def test_chat_generic_exception(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
|
||||
api = module.TrialChatApi()
|
||||
@ -525,17 +549,17 @@ class TestTrialChatApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(InternalServerError):
|
||||
method(api, MagicMock(), account, trial_app_chat)
|
||||
method(api, self.sqlite_session, account, trial_app_chat)
|
||||
|
||||
|
||||
class TestTrialCompletionApi:
|
||||
class TestTrialCompletionApi(_UsesSQLiteSession):
|
||||
def test_not_completion_app(self, app: Flask, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
method = unwrap(api.post)
|
||||
|
||||
with app.test_request_context("/", json={"inputs": {}, "query": ""}):
|
||||
with pytest.raises(NotCompletionAppError):
|
||||
method(api, MagicMock(), account, MagicMock(mode=AppMode.CHAT))
|
||||
method(api, self.sqlite_session, account, MagicMock(mode=AppMode.CHAT))
|
||||
|
||||
def test_success(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -546,7 +570,7 @@ class TestTrialCompletionApi:
|
||||
patch.object(module.AppGenerateService, "generate", return_value=MagicMock()),
|
||||
patch.object(module.RecommendedAppService, "add_trial_app_record"),
|
||||
):
|
||||
result = method(api, MagicMock(), account, trial_app_completion)
|
||||
result = method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
assert result is not None
|
||||
|
||||
@ -563,7 +587,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(AppUnavailableError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
def test_completion_provider_not_init(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -578,7 +602,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderNotInitializeError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
def test_completion_quota_exceeded(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -593,7 +617,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderQuotaExceededError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
def test_completion_model_not_support(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -608,7 +632,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProviderModelCurrentlyNotSupportError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
def test_completion_invoke_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -623,7 +647,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(CompletionRequestError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
def test_completion_rate_limit_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -638,7 +662,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(InternalServerError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
def test_completion_value_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -653,7 +677,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ValueError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
def test_completion_generic_exception(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
|
||||
api = module.TrialCompletionApi()
|
||||
@ -668,7 +692,7 @@ class TestTrialCompletionApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(InternalServerError):
|
||||
method(api, MagicMock(), account, trial_app_completion)
|
||||
method(api, self.sqlite_session, account, trial_app_completion)
|
||||
|
||||
|
||||
class TestTrialMessageSuggestedQuestionApi:
|
||||
@ -1131,49 +1155,55 @@ class TestTrialAppWorkflowTaskStopApi:
|
||||
|
||||
|
||||
class TestTrialSitApi:
|
||||
def test_no_site(self, app: Flask) -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(Site,)], indirect=True)
|
||||
def test_no_site(
|
||||
self,
|
||||
app: Flask,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
api = module.TrialSitApi()
|
||||
method = unwrap(api.get)
|
||||
app_model = MagicMock()
|
||||
app_model.id = "a1"
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = None
|
||||
app_model.id = str(uuid4())
|
||||
|
||||
with app.test_request_context("/"):
|
||||
with pytest.raises(Forbidden):
|
||||
method(api, session, app_model)
|
||||
method(api, sqlite_session, app_model)
|
||||
|
||||
session.scalar.assert_called_once()
|
||||
|
||||
def test_archived_tenant(self, app: Flask) -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(Site,)], indirect=True)
|
||||
def test_archived_tenant(
|
||||
self,
|
||||
app: Flask,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
api = module.TrialSitApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
site = MagicMock()
|
||||
app_model = SimpleNamespace(id="a1", tenant_id="tenant-1")
|
||||
app_model = SimpleNamespace(id=str(uuid4()), tenant_id="tenant-1")
|
||||
tenant = SimpleNamespace(status=TenantStatus.ARCHIVE)
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = site
|
||||
_persist_site(sqlite_session, app_model.id)
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(module.TenantService, "get_tenant_by_id", return_value=tenant) as get_tenant_by_id,
|
||||
):
|
||||
with pytest.raises(Forbidden):
|
||||
method(api, session, app_model)
|
||||
method(api, sqlite_session, app_model)
|
||||
|
||||
session.scalar.assert_called_once()
|
||||
get_tenant_by_id.assert_called_once_with("tenant-1", session=session)
|
||||
get_tenant_by_id.assert_called_once_with("tenant-1", session=sqlite_session)
|
||||
|
||||
def test_success(self, app: Flask) -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(Site,)], indirect=True)
|
||||
def test_success(
|
||||
self,
|
||||
app: Flask,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
api = module.TrialSitApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
site = MagicMock()
|
||||
app_model = SimpleNamespace(id="a1", tenant_id="tenant-1")
|
||||
app_model = SimpleNamespace(id=str(uuid4()), tenant_id="tenant-1")
|
||||
tenant = SimpleNamespace(status=TenantStatus.NORMAL)
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = site
|
||||
site = _persist_site(sqlite_session, app_model.id)
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
@ -1183,11 +1213,11 @@ class TestTrialSitApi:
|
||||
mock_validate_result = MagicMock()
|
||||
mock_validate_result.model_dump.return_value = {"name": "test", "icon": "icon"}
|
||||
mock_validate.return_value = mock_validate_result
|
||||
result = method(api, session, app_model)
|
||||
result = method(api, sqlite_session, app_model)
|
||||
|
||||
assert result == {"name": "test", "icon": "icon"}
|
||||
session.scalar.assert_called_once()
|
||||
get_tenant_by_id.assert_called_once_with("tenant-1", session=session)
|
||||
get_tenant_by_id.assert_called_once_with("tenant-1", session=sqlite_session)
|
||||
mock_validate.assert_called_once_with(site)
|
||||
|
||||
|
||||
class TestAppWorkflowApi:
|
||||
|
||||
Loading…
Reference in New Issue
Block a user