test: use sqlite3 session in test_trial (#38732)

This commit is contained in:
Asuka Minato 2026-07-30 13:47:18 +09:00 committed by GitHub
parent 814fee5550
commit b852252076
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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: