diff --git a/api/tests/unit_tests/controllers/console/explore/test_trial.py b/api/tests/unit_tests/controllers/console/explore/test_trial.py index 5c2d20c75a1..1df0cec37e6 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -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: