From 38a8ac6f462cfe2853f073bc1d5f48ff316374a3 Mon Sep 17 00:00:00 2001 From: hjlarry Date: Mon, 21 Sep 2026 14:59:31 +0800 Subject: [PATCH] refactor(api): isolate trial preview queries and admission --- api/.importlinter | 14 + .../console/app/preview_admission.py | 26 ++ api/controllers/console/explore/error.py | 12 + api/controllers/console/explore/trial.py | 80 ++-- api/extensions/ext_application_services.py | 18 +- .../app_definition_query_repository.py | 6 +- .../app_preview_query_repository.py | 79 ++++ api/services/app_preview_query_service.py | 82 ++++ .../controllers/console/explore/test_trial.py | 173 +------ .../console/explore/test_trial_preview.py | 422 ++++++++++++++++++ .../test_ext_application_services.py | 28 ++ .../test_app_preview_query_repository.py | 342 ++++++++++++++ 12 files changed, 1068 insertions(+), 214 deletions(-) create mode 100644 api/controllers/console/app/preview_admission.py create mode 100644 api/repositories/app_preview_query_repository.py create mode 100644 api/services/app_preview_query_service.py create mode 100644 api/tests/unit_tests/controllers/console/explore/test_trial_preview.py create mode 100644 api/tests/unit_tests/repositories/test_app_preview_query_repository.py diff --git a/api/.importlinter b/api/.importlinter index 33c4afdedd4658..4f3726960e8ca5 100644 --- a/api/.importlinter +++ b/api/.importlinter @@ -584,6 +584,20 @@ forbidden_modules = sqlalchemy werkzeug +[importlinter:contract:app-preview-query-boundary] +name = App preview queries are framework and persistence neutral +type = forbidden +source_modules = + services.app_preview_query_service +forbidden_modules = + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + [importlinter:contract:message-suggested-questions-boundary] name = Message suggested questions port is framework and persistence neutral type = forbidden diff --git a/api/controllers/console/app/preview_admission.py b/api/controllers/console/app/preview_admission.py new file mode 100644 index 00000000000000..e799131c8c0bd2 --- /dev/null +++ b/api/controllers/console/app/preview_admission.py @@ -0,0 +1,26 @@ +"""Public catalog preview admission, separate from authenticated trial execution.""" + +from collections.abc import Callable +from functools import wraps +from typing import Concatenate + +from controllers.console.app.error import AppNotFoundError +from extensions.ext_application_services import application_services +from services.app_preview_query_service import AppPreviewRef, AppPreviewUnavailableError + + +def get_preview_app[T, **P, R]( + view: Callable[Concatenate[T, AppPreviewRef, P], R], +) -> Callable[Concatenate[T, P], R]: + @wraps(view) + def decorated(self: T, /, *args: P.args, **kwargs: P.kwargs) -> R: + app_id = kwargs.pop("app_id", None) + if app_id is None: + raise RuntimeError("The app preview admission route must provide app_id") + try: + app = application_services().app_previews.get_access(app_id=str(app_id)) + except AppPreviewUnavailableError as error: + raise AppNotFoundError() from error + return view(self, app, *args, **kwargs) + + return decorated diff --git a/api/controllers/console/explore/error.py b/api/controllers/console/explore/error.py index 6f0ba500b8465b..c8a52d640b7517 100644 --- a/api/controllers/console/explore/error.py +++ b/api/controllers/console/explore/error.py @@ -37,6 +37,18 @@ class RecommendedAppNotFoundError(BaseHTTPException): code = 404 +class AppPreviewSiteUnavailableError(BaseHTTPException): + error_code = "app_site_unavailable" + description = "The app preview site is unavailable." + code = 403 + + +class AppPreviewOwnerUnavailableError(BaseHTTPException): + error_code = "app_owner_unavailable" + description = "The app preview owner is unavailable." + code = 403 + + class TrialAppNotAllowed(BaseHTTPException): """*403* `Trial App Not Allowed` diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index f4334229ba3e17..896ba86f5f2f2a 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -7,9 +7,8 @@ from flask import Response, request from flask_restx import Resource from pydantic import AliasChoices, BaseModel, Field, field_validator -from sqlalchemy import select from sqlalchemy.orm import Session -from werkzeug.exceptions import Forbidden, InternalServerError, NotFound, Unauthorized +from werkzeug.exceptions import InternalServerError, NotFound, Unauthorized import services from controllers.common.fields import ( @@ -27,6 +26,7 @@ ) from controllers.console import console_ns from controllers.console.app.error import ( + AppNotFoundError, AppUnavailableError, AudioTooLargeError, CompletionRequestError, @@ -40,7 +40,14 @@ SpeechToTextDisabledError, UnsupportedAudioTypeError, ) +from controllers.console.app.preview_admission import get_preview_app from controllers.console.app.wraps import get_previewable_app_model, with_session +from controllers.console.explore.error import ( + AppPreviewOwnerUnavailableError as AppPreviewOwnerUnavailableHttpError, +) +from controllers.console.explore.error import ( + AppPreviewSiteUnavailableError as AppPreviewSiteUnavailableHttpError, +) from controllers.console.explore.error import ( AppSuggestedQuestionsAfterAnswerDisabledError, NotChatAppError, @@ -74,17 +81,19 @@ from libs.helper import dump_response, to_timestamp, uuid_value from machinery.context import RequestContext from models import Account -from models.account import TenantStatus from models.enums import CreatorUserRole -from models.model import Site from models.workflow import Workflow from services.account_errors import AccountNotFoundError -from services.account_service import TenantService from services.app_definition_query_service import AppDefinitionUnavailableError +from services.app_preview_query_service import ( + AppPreviewOwnerUnavailableError, + AppPreviewRef, + AppPreviewSiteUnavailableError, + AppPreviewUnavailableError, +) from services.app_ref_service import AppRefService from services.app_service import AppResponseView, AppService from services.audio_service import AudioService -from services.dataset_service import DatasetService from services.errors.audio import ( AudioTooLargeServiceError, NoAudioUploadedServiceError, @@ -807,40 +816,31 @@ class TrialSitApi(Resource): """Resource for trial app sites.""" @console_ns.response(200, "Success", console_ns.models[SiteResponse.__name__]) - @with_session(write=False) - @get_previewable_app_model(None) - def get(self, session: Session, app_model): + @get_preview_app + def get(self, app: AppPreviewRef) -> dict[str, object]: """Retrieve app site info. Returns the site configuration for the application including theme, icons, and text. """ - site = session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) - - if not site: - raise Forbidden() - - tenant = TenantService.get_tenant_by_id(app_model.tenant_id, session=session) - assert tenant - if tenant.status == TenantStatus.ARCHIVE: - raise Forbidden() - - return SiteResponse.model_validate(site).model_dump(mode="json") + try: + site = application_services().app_previews.get_site(app=app) + except AppPreviewSiteUnavailableError as error: + raise AppPreviewSiteUnavailableHttpError(str(error)) from error + except AppPreviewOwnerUnavailableError as error: + raise AppPreviewOwnerUnavailableHttpError(str(error)) from error + return dump_response(SiteResponse, site) class TrialAppParameterApi(Resource): """Resource for app variables.""" @console_ns.response(200, "Success", console_ns.models[ParametersResponse.__name__]) - @with_session(write=False) - @get_previewable_app_model(None) - def get(self, session: Session, app_model): + @get_preview_app + def get(self, app: AppPreviewRef) -> dict[str, object]: """Retrieve app parameters.""" - if app_model is None: - raise AppUnavailableError() - try: - parameters = application_services().app_definitions.get_parameters(app_model.id) + parameters = application_services().app_definitions.get_parameters(app.app_id) except AppDefinitionUnavailableError: raise AppUnavailableError() from None @@ -885,20 +885,28 @@ def get(self, session: Session, app_model): class DatasetListApi(Resource): @console_ns.doc(params=query_params_from_model(TrialDatasetListQuery)) @console_ns.response(200, "Success", console_ns.models[TrialDatasetListResponse.__name__]) - @with_session(write=False) - @get_previewable_app_model(None) - def get(self, session: Session, app_model): + @get_preview_app + def get(self, app: AppPreviewRef) -> dict[str, object]: + # These legacy fields are response metadata: the query returns all + # requested IDs. Keep their integer fallback and echo behavior. page = request.args.get("page", default=1, type=int) limit = request.args.get("limit", default=20, type=int) ids = request.args.getlist("ids") - tenant_id = app_model.tenant_id - if ids: - datasets, total = DatasetService.get_datasets_by_ids(ids, tenant_id, session=session) - else: + if not ids: raise NeedAddIdsError() - - response = {"data": datasets, "has_more": len(datasets) == limit, "limit": limit, "total": total, "page": page} + try: + datasets = application_services().app_previews.get_datasets(app=app, ids=ids) + except AppPreviewUnavailableError as error: + raise AppNotFoundError() from error + + response = { + "data": datasets, + "has_more": len(datasets) == limit, + "limit": limit, + "total": len(datasets), + "page": page, + } return dump_response(TrialDatasetListResponse, response) diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index 64595b1a1ef4e1..94291ce7939124 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -39,6 +39,7 @@ ) from repositories.account_repository import SQLAlchemyAccountRepository from repositories.app_definition_query_repository import AppDefinitionQueryRepository +from repositories.app_preview_query_repository import AppPreviewQueryRepository from repositories.app_site_command_repository import AppSiteCommandRepository from repositories.app_statistic_query_repository import AppStatisticQueryRepository from repositories.app_tracing_config_repository import SQLAlchemyAppTracingConfigRepository @@ -136,6 +137,7 @@ from services.account_password_service import AccountPasswordService from services.account_profile_service import AccountProfileService from services.app_definition_query_service import AppDefinitionQueryService +from services.app_preview_query_service import AppPreviewQueryService from services.app_site_service import AppSiteService from services.app_statistic_query import AppStatisticQuery from services.app_task_service import AppTaskControlService @@ -260,6 +262,7 @@ class ApplicationServices: accounts: AccountServices account_activation: AccountActivationService app_definitions: AppDefinitionQueryService + app_previews: AppPreviewQueryService app_sites: AppSiteService app_statistics: AppStatisticQuery app_tracing_configs: AppTracingConfigService @@ -446,6 +449,11 @@ def build_application_services( database=database_catalog, builtin=builtin_catalog, ) + recommended_app_queries = RecommendedAppQueryService( + catalog=recommended_app_catalog, + trial_apps=trial_apps, + trial_enabled=trial_app_enabled, + ) workspace_query_repository = WorkspaceQueryRepository(session_factory=database_client) file_service = FileService(session_factory=database_client) remote_file_service = RemoteFileService(files=file_service) @@ -618,6 +626,10 @@ def build_application_services( dify_config.CONSOLE_API_URL + "/console/api/workspaces/current/tool-provider/builtin/" ), ), + app_previews=AppPreviewQueryService( + apps=AppPreviewQueryRepository(session_factory=database_client), + is_previewable=recommended_app_queries.is_previewable, + ), app_sites=AppSiteService( sites=AppSiteCommandRepository(session_factory=database_client), ), @@ -716,11 +728,7 @@ def build_application_services( partner_tenant_bindings=PartnerTenantBindingService( sync_bindings=BillingService.sync_partner_tenants_bindings, ), - recommended_app_queries=RecommendedAppQueryService( - catalog=recommended_app_catalog, - trial_apps=trial_apps, - trial_enabled=trial_app_enabled, - ), + recommended_app_queries=recommended_app_queries, remote_files=remote_file_service, app_tasks=AppTaskControlService(redis_client=redis), trial_app_access=TrialAppAccessService(apps=trial_apps), diff --git a/api/repositories/app_definition_query_repository.py b/api/repositories/app_definition_query_repository.py index 0907040e787eb4..eeaecd55d32049 100644 --- a/api/repositories/app_definition_query_repository.py +++ b/api/repositories/app_definition_query_repository.py @@ -25,7 +25,7 @@ from services.web_app_runtime_query_service import WebAppRuntimeRecord -def _map_site_configuration(site: Site) -> AppSiteConfiguration: +def map_site_configuration(site: Site) -> AppSiteConfiguration: return AppSiteConfiguration( title=site.title, chat_color_theme=site.chat_color_theme, @@ -175,7 +175,7 @@ def get_site_configuration(self, app_id: str) -> AppSiteConfiguration | None: if site is None: return None - return _map_site_configuration(site) + return map_site_configuration(site) def get_runtime_record(self, app_id: str) -> WebAppRuntimeRecord | None: with self._session_factory() as session: @@ -194,7 +194,7 @@ def get_runtime_record(self, app_id: str) -> WebAppRuntimeRecord | None: app_id = app.id tenant_id = app.tenant_id enable_site = app.enable_site - site_configuration = _map_site_configuration(site) + site_configuration = map_site_configuration(site) plan = tenant.plan tenant_status = tenant.status.value tenant_custom_config_json = tenant.custom_config diff --git a/api/repositories/app_preview_query_repository.py b/api/repositories/app_preview_query_repository.py new file mode 100644 index 00000000000000..9a08b236327eef --- /dev/null +++ b/api/repositories/app_preview_query_repository.py @@ -0,0 +1,79 @@ +"""Load detached app-preview data without trial or account admission policy.""" + +from collections.abc import Sequence +from typing import override + +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Tenant +from models.dataset import Dataset +from models.model import App, Site +from repositories.app_definition_query_repository import map_site_configuration +from services.app_preview_query_service import ( + AppPreviewDataset, + AppPreviewQuery, + AppPreviewRef, + AppPreviewSite, +) + + +class AppPreviewQueryRepository(AppPreviewQuery): + def __init__(self, *, session_factory: sessionmaker[Session]) -> None: + self._session_factory: sessionmaker[Session] = session_factory + + @override + def get_app(self, *, app_id: str) -> AppPreviewRef | None: + with self._session_factory() as session: + row = session.execute( + select(App.id, App.tenant_id).where(App.id == app_id, App.status == "normal") + ).one_or_none() + if row is None: + return None + return AppPreviewRef(app_id=row.id, tenant_id=row.tenant_id) + + @override + def get_site(self, *, app: AppPreviewRef) -> AppPreviewSite | None: + with self._session_factory() as session: + row = session.execute( + select(Site, Tenant.status) + .select_from(App) + .join(Site, Site.app_id == App.id) + .outerjoin(Tenant, Tenant.id == App.tenant_id) + .where(App.id == app.app_id, App.tenant_id == app.tenant_id, App.status == "normal") + .limit(1) + ).first() + if row is None: + return None + site, owner_status = row + return AppPreviewSite( + configuration=map_site_configuration(site), + owner_status=owner_status.value if owner_status is not None else None, + ) + + @override + def get_datasets(self, *, app: AppPreviewRef, ids: Sequence[str]) -> tuple[AppPreviewDataset, ...] | None: + with self._session_factory() as session: + app_id = session.scalar( + select(App.id).where(App.id == app.app_id, App.tenant_id == app.tenant_id, App.status == "normal") + ) + if app_id is None: + return None + if not ids: + return () + datasets = session.scalars(select(Dataset).where(Dataset.id.in_(ids), Dataset.tenant_id == app.tenant_id)) + return tuple( + AppPreviewDataset( + id=dataset.id, + name=dataset.name, + description=dataset.description, + permission=dataset.permission.value if dataset.permission is not None else None, + data_source_type=dataset.data_source_type.value if dataset.data_source_type is not None else None, + indexing_technique=dataset.indexing_technique.value + if dataset.indexing_technique is not None + else None, + created_by=dataset.created_by, + created_at=dataset.created_at, + ) + for dataset in datasets + ) diff --git a/api/services/app_preview_query_service.py b/api/services/app_preview_query_service.py new file mode 100644 index 00000000000000..6d4f8406f480e3 --- /dev/null +++ b/api/services/app_preview_query_service.py @@ -0,0 +1,82 @@ +"""Read-only app previews, independent of trial execution and account quotas.""" + +from collections.abc import Callable, Sequence +from dataclasses import dataclass +from datetime import datetime +from typing import Protocol + +from services.app_definition_query_service import AppSiteConfiguration + + +@dataclass(frozen=True, slots=True) +class AppPreviewRef: + app_id: str + tenant_id: str + + +@dataclass(frozen=True, slots=True) +class AppPreviewSite: + configuration: AppSiteConfiguration + owner_status: str | None + + +@dataclass(frozen=True, slots=True) +class AppPreviewDataset: + id: str + name: str + description: str | None + permission: str | None + data_source_type: str | None + indexing_technique: str | None + created_by: str | None + created_at: datetime | None + + +class AppPreviewQuery(Protocol): + def get_app(self, *, app_id: str) -> AppPreviewRef | None: ... + + def get_site(self, *, app: AppPreviewRef) -> AppPreviewSite | None: ... + + def get_datasets(self, *, app: AppPreviewRef, ids: Sequence[str]) -> tuple[AppPreviewDataset, ...] | None: ... + + +class AppPreviewUnavailableError(LookupError): + """The app is not in the preview catalog or is no longer available.""" + + +class AppPreviewSiteUnavailableError(LookupError): + """The admitted app no longer has an available site.""" + + +class AppPreviewOwnerUnavailableError(LookupError): + """The site owner is missing or archived.""" + + +class AppPreviewQueryService: + def __init__(self, *, apps: AppPreviewQuery, is_previewable: Callable[[str], bool]) -> None: + self._apps: AppPreviewQuery = apps + self._is_previewable: Callable[[str], bool] = is_previewable + + def get_access(self, *, app_id: str) -> AppPreviewRef: + # Catalog membership can involve remote I/O. Resolve it before opening + # the app query session, without imposing trial execution admission. + if not self._is_previewable(app_id): + raise AppPreviewUnavailableError(f"App {app_id} is not available for preview") + app = self._apps.get_app(app_id=app_id) + if app is None: + raise AppPreviewUnavailableError(f"App {app_id} is not available for preview") + return app + + def get_site(self, *, app: AppPreviewRef) -> AppSiteConfiguration: + site = self._apps.get_site(app=app) + if site is None: + raise AppPreviewSiteUnavailableError(f"Site for app {app.app_id} is unavailable") + if site.owner_status is None or site.owner_status == "archive": + raise AppPreviewOwnerUnavailableError(f"Owner of app {app.app_id} is unavailable") + return site.configuration + + def get_datasets(self, *, app: AppPreviewRef, ids: Sequence[str]) -> tuple[AppPreviewDataset, ...]: + datasets = self._apps.get_datasets(app=app, ids=ids) + if datasets is None: + raise AppPreviewUnavailableError(f"App {app.app_id} is no longer available in tenant {app.tenant_id}") + return datasets 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 a91c8e91ad5d3b..650743004932a6 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -5,16 +5,13 @@ from types import SimpleNamespace from typing import Any from unittest.mock import MagicMock, patch -from uuid import uuid4 import pytest from flask import Flask, request from sqlalchemy.orm import Session -from werkzeug.exceptions import Forbidden import controllers.console.explore.trial as module from controllers.console.app.error import ( - AppUnavailableError, CompletionRequestError, ProviderModelCurrentlyNotSupportError, ProviderNotInitializeError, @@ -22,7 +19,6 @@ SpeechToTextDisabledError, ) from controllers.console.explore.trial import TextToSpeechRequest -from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.errors.error import ( ModelCurrentlyNotSupportError, ProviderTokenNotInitError, @@ -32,10 +28,8 @@ from core.workflow.llm_environment_variable import LLMEnvironmentVariable from graphon.model_runtime.errors.invoke import InvokeError from graphon.variables import SecretVariable, StringVariable -from models import Account, Tenant -from models.account import TenantStatus -from models.dataset import Dataset -from models.model import App, AppMode, Site +from models import Account +from models.model import App, AppMode from models.tools import WorkflowToolProvider from models.workflow import Workflow from services.app_ref_service import AppRef, MessageRef @@ -79,18 +73,6 @@ 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() -> App: return _app(app_id="a-chat", mode=AppMode.CHAT) @@ -101,57 +83,9 @@ def test_trial_workflow_uses_trial_scoped_simple_account_model() -> None: assert module.simple_account_model.__schema__["properties"].keys() >= {"id", "name", "email"} -def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask, unbound_session: Session): - api = module.DatasetListApi() - method = unwrap(api.get) - app_model = _app(app_id="app-1", mode=AppMode.CHAT) - dataset = Dataset( - id="dataset-1", - tenant_id=app_model.tenant_id, - name="Dataset", - description="description", - permission="only_me", - data_source_type="upload_file", - indexing_technique="high_quality", - created_by="user-1", - created_at=datetime(2024, 1, 1, tzinfo=UTC), - ) - dataset.permission_keys = ["dataset.acl.readonly"] # type: ignore[attr-defined] - with ( - app.test_request_context("/?page=1&limit=20&ids=dataset-1"), - patch.object( - module.DatasetService, - "get_datasets_by_ids", - return_value=([dataset], 1), - ) as get_datasets, - ): - result = method(api, unbound_session, app_model) - - get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=unbound_session) - assert result == { - "data": [ - { - "id": "dataset-1", - "name": "Dataset", - "description": "description", - "permission": "only_me", - "data_source_type": "upload_file", - "indexing_technique": "high_quality", - "created_by": "user-1", - "created_at": 1704067200, - "permission_keys": ["dataset.acl.readonly"], - } - ], - "has_more": False, - "limit": 20, - "total": 1, - "page": 1, - } - - @pytest.mark.parametrize( "api_type", - [module.TrialSitApi, module.TrialAppParameterApi, module.AppApi, module.AppWorkflowApi, module.DatasetListApi], + [module.AppApi, module.AppWorkflowApi], ) def test_preview_handlers_use_explicit_read_session(api_type: type) -> None: source = getsource(api_type.get) @@ -182,43 +116,6 @@ def test_trial_app_detail_serializes_with_explicit_session( module.TrialAppDetailResponse.model_validate.assert_called_once_with(response_view, from_attributes=True) -class TestTrialAppParameterApi: - def test_app_unavailable(self, unbound_session: Session) -> None: - api = module.TrialAppParameterApi() - method = unwrap(api.get) - - with pytest.raises(AppUnavailableError): - method(api, unbound_session, None) - - def test_success(self, unbound_session: Session) -> None: - api = module.TrialAppParameterApi() - method = unwrap(api.get) - parameters = get_parameters_from_feature_dict(features_dict={}, user_input_form=[]) - expected = module.ParametersResponse.model_validate(parameters).model_dump(mode="json") - app_definitions = MagicMock() - app_definitions.get_parameters.return_value = parameters - services = SimpleNamespace(app_definitions=app_definitions) - - with patch.object(module, "application_services", return_value=services): - result = method(api, unbound_session, _app(app_id="app-1", mode=AppMode.CHAT)) - - assert result == expected - app_definitions.get_parameters.assert_called_once_with("app-1") - - def test_unavailable_parameters(self, unbound_session: Session) -> None: - api = module.TrialAppParameterApi() - method = unwrap(api.get) - app_definitions = MagicMock() - app_definitions.get_parameters.side_effect = module.AppDefinitionUnavailableError - services = SimpleNamespace(app_definitions=app_definitions) - - with ( - patch.object(module, "application_services", return_value=services), - pytest.raises(AppUnavailableError), - ): - method(api, unbound_session, _app(app_id="app-1", mode=AppMode.CHAT)) - - class TestTrialChatAudioApi: def test_success( self, @@ -641,70 +538,6 @@ def test_invoke_error(self, app: Flask, trial_app_chat: App, account: Account) - ) -class TestTrialSitApi: - def test_no_site( - self, - app: Flask, - sqlite_session: Session, - ) -> None: - api = module.TrialSitApi() - method = unwrap(api.get) - app_model = _app(app_id=str(uuid4()), mode=AppMode.CHAT) - - with app.test_request_context("/"): - with pytest.raises(Forbidden): - method(api, sqlite_session, app_model) - - def test_archived_tenant( - self, - app: Flask, - sqlite_session: Session, - ) -> None: - api = module.TrialSitApi() - method = unwrap(api.get) - - app_model = _app(app_id=str(uuid4()), mode=AppMode.CHAT) - tenant = Tenant(name="Archived Tenant", status=TenantStatus.ARCHIVE) - tenant.id = app_model.tenant_id - _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, sqlite_session, app_model) - - get_tenant_by_id.assert_called_once_with("tenant-1", session=sqlite_session) - - def test_success( - self, - app: Flask, - sqlite_session: Session, - ) -> None: - api = module.TrialSitApi() - method = unwrap(api.get) - - app_model = _app(app_id=str(uuid4()), mode=AppMode.CHAT) - tenant = Tenant(name="Active Tenant", status=TenantStatus.NORMAL) - tenant.id = app_model.tenant_id - 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, - patch.object(module.SiteResponse, "model_validate") as mock_validate, - ): - 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, sqlite_session, app_model) - - assert result == {"name": "test", "icon": "icon"} - get_tenant_by_id.assert_called_once_with("tenant-1", session=sqlite_session) - mock_validate.assert_called_once_with(site) - - class TestAppWorkflowApi: def test_uses_injected_session(self, sqlite_session: Session) -> None: api = module.AppWorkflowApi() diff --git a/api/tests/unit_tests/controllers/console/explore/test_trial_preview.py b/api/tests/unit_tests/controllers/console/explore/test_trial_preview.py new file mode 100644 index 00000000000000..69fa03cf921c60 --- /dev/null +++ b/api/tests/unit_tests/controllers/console/explore/test_trial_preview.py @@ -0,0 +1,422 @@ +"""Anonymous app previews through HTTP and real SQLite query boundaries.""" + +import json +from collections.abc import Callable +from dataclasses import dataclass, field +from datetime import datetime +from typing import Literal +from uuid import uuid4 + +import pytest +from flask import Flask +from sqlalchemy import Connection, Engine, delete, event, select, text, update +from sqlalchemy.orm import Session, SessionTransaction, sessionmaker +from sqlalchemy.pool import QueuePool +from werkzeug.test import TestResponse + +import controllers.common.fields as fields_module +import controllers.console.explore.trial as trial_module +import controllers.console.wraps as console_wraps +import libs.login as login_module +from controllers.console.app import preview_admission as admission_module +from libs.external_api import ExternalApi +from models import AccountTrialAppRecord, App, AppMode, Tenant, TrialApp +from models.account import TenantStatus +from models.dataset import Dataset +from models.enums import CustomizeTokenStrategy +from models.model import AppModelConfig, IconType, Site +from repositories.app_definition_query_repository import AppDefinitionQueryRepository +from repositories.app_preview_query_repository import AppPreviewQueryRepository +from repositories.trial_app_repository import TrialAppRepository +from services.app_definition_query_service import AppDefinitionQueryService +from services.app_preview_query_service import AppPreviewQueryService +from services.recommended_app_query_service import ( + RecommendedAppCatalogPage, + RecommendedAppDetailRecord, + RecommendedAppQueryService, +) + +_Endpoint = Literal["parameters", "site", "datasets"] +_CREATED_AT = datetime(2024, 1, 1) +_INPUT_FORM = [{"number": {"label": "Count", "variable": "count", "required": False, "default": 0}}] + + +@dataclass +class _Catalog: + engine: Engine + ids: set[str] = field(default_factory=set) + calls: list[str] = field(default_factory=list) + + def contains(self, app_id: str) -> bool: + assert isinstance(self.engine.pool, QueuePool) + assert self.engine.pool.checkedout() == 0 + self.calls.append(app_id) + return app_id in self.ids + + def list_recommended(self, language: str) -> RecommendedAppCatalogPage: + raise AssertionError(f"Unexpected list request for {language}") + + def list_learn_dify(self, language: str) -> RecommendedAppCatalogPage: + raise AssertionError(f"Unexpected list request for {language}") + + def get_detail(self, app_id: str) -> RecommendedAppDetailRecord | None: + raise AssertionError(f"Unexpected detail request for {app_id}") + + +@dataclass(frozen=True) +class _ApplicationServices: + app_previews: AppPreviewQueryService + app_definitions: AppDefinitionQueryService + recommended_app_queries: RecommendedAppQueryService + + +@dataclass(frozen=True) +class _Harness: + app: Flask + owner: Tenant + target: App + listing: TrialApp + site: Site + config: AppModelConfig + dataset: Dataset + creator_id: str + factory: sessionmaker[Session] + engine: Engine + catalog: _Catalog + sessions: list[Session] + signed_icons: list[str] + + def get(self, endpoint: _Endpoint, *, query: str = "", app_id: str | None = None) -> TestResponse: + response = self.app.test_client().get(f"/trial-apps/{app_id or self.target.id}/{endpoint}{query}") + assert response.headers["Content-Type"] == "application/json" + assert int(response.headers["Content-Length"]) == len(response.data) + assert isinstance(self.engine.pool, QueuePool) + assert self.engine.pool.checkedout() == 0 + assert all(not session.in_transaction() and not session.identity_map for session in self.sessions) + return response + + def dataset_query(self) -> str: + return f"?ids={self.dataset.id}" + + +@pytest.fixture +def harness( + monkeypatch: pytest.MonkeyPatch, + config_overrides: Callable[..., None], + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], +) -> _Harness: + config_overrides( + LOGIN_DISABLED=False, + INIT_PASSWORD="preview-does-not-require-setup", + UPLOAD_IMAGE_FILE_SIZE_LIMIT=1, + UPLOAD_VIDEO_FILE_SIZE_LIMIT=2, + UPLOAD_AUDIO_FILE_SIZE_LIMIT=3, + UPLOAD_FILE_SIZE_LIMIT=4, + WORKFLOW_FILE_UPLOAD_LIMIT=5, + ) + owner = Tenant(name="App owner") + creator_id = str(uuid4()) + target = App( + id=str(uuid4()), tenant_id=owner.id, name="Preview", mode=AppMode.CHAT, enable_site=False, enable_api=False + ) + listing = TrialApp(app_id=target.id, tenant_id=str(uuid4()), trial_limit=0) + config = AppModelConfig( + app_id=target.id, + opening_statement="", + suggested_questions="[]", + suggested_questions_after_answer='{"enabled":false}', + speech_to_text='{"enabled":false}', + text_to_speech='{"enabled":false,"voice":"","autoPlay":"disabled"}', + retriever_resource='{"enabled":false}', + more_like_this='{"enabled":false}', + sensitive_word_avoidance='{"enabled":false,"type":"","configs":[]}', + file_upload='{"enabled":false,"number_limits":0}', + user_input_form=json.dumps(_INPUT_FORM), + ) + site = Site( + app_id=target.id, + title="Preview site", + default_language="en-US", + customize_token_strategy=CustomizeTokenStrategy.UUID, + icon_type=IconType.IMAGE, + icon=str(uuid4()), + icon_background="", + chat_color_theme="", + chat_color_theme_inverted=False, + description=None, + copyright="", + privacy_policy=None, + input_placeholder="", + custom_disclaimer="", + show_workflow_steps=False, + use_icon_as_answer_icon=False, + prompt_public=False, + ) + dataset = Dataset( + tenant_id=owner.id, + name="Owner dataset", + description=None, + permission="only_me", + data_source_type="upload_file", + indexing_technique=None, + created_by=creator_id, + created_at=_CREATED_AT, + ) + with sqlite_session_factory.begin() as session: + session.add_all([owner, target, listing, config, site, dataset]) + session.flush() + target.app_model_config_id = config.id + session.add(AccountTrialAppRecord(app_id=target.id, account_id=creator_id, count=100)) + + factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + sessions: list[Session] = [] + + @event.listens_for(factory, "after_begin") + def track_session(session: Session, _transaction: SessionTransaction, _connection: Connection) -> None: + assert all(not previous.in_transaction() for previous in sessions) + assert isinstance(sqlite_engine.pool, QueuePool) + assert sqlite_engine.pool.checkedout() == 1 + sessions.append(session) + + catalog = _Catalog(sqlite_engine) + recommendations = RecommendedAppQueryService( + catalog=catalog, trial_apps=TrialAppRepository(factory), trial_enabled=False + ) + services = _ApplicationServices( + app_previews=AppPreviewQueryService( + apps=AppPreviewQueryRepository(session_factory=factory), is_previewable=recommendations.is_previewable + ), + app_definitions=AppDefinitionQueryService( + definitions=AppDefinitionQueryRepository(session_factory=factory), builtin_icon_url_prefix="/tools/" + ), + recommended_app_queries=recommendations, + ) + signed_icons: list[str] = [] + + def sign_icon(file_id: str) -> str: + assert isinstance(sqlite_engine.pool, QueuePool) + assert sqlite_engine.pool.checkedout() == 0 + signed_icons.append(file_id) + return f"https://files.example/{file_id}?sign=preview" + + for module in (trial_module, admission_module): + monkeypatch.setattr(module, "application_services", lambda: services) + monkeypatch.setattr(login_module, "current_user", None) + monkeypatch.setattr(console_wraps, "_is_setup_completed", lambda: False) + monkeypatch.setattr(fields_module.file_helpers, "get_signed_file_url", sign_icon) + app = Flask(__name__) + app.config.update(TESTING=True, RESTX_ERROR_404_HELP=False) + api = ExternalApi(app) + api.add_resource(trial_module.TrialAppParameterApi, "/trial-apps//parameters") + api.add_resource(trial_module.TrialSitApi, "/trial-apps//site") + api.add_resource(trial_module.DatasetListApi, "/trial-apps//datasets") + return _Harness( + app, + owner, + target, + listing, + site, + config, + dataset, + creator_id, + sqlite_session_factory, + sqlite_engine, + catalog, + sessions, + signed_icons, + ) + + +def _assert_error(response: TestResponse, status: int, code: str) -> None: + assert response.status_code == status, response.get_json() + body = response.get_json() + assert body["code"] == code + assert body["status"] == status + assert isinstance(body["message"], str) + assert body["message"] + + +@pytest.mark.parametrize("endpoint", ["parameters", "site", "datasets"]) +@pytest.mark.parametrize("membership", ["trial", "catalog"]) +def test_anonymous_preview_ignores_trial_feature_quota_and_setup( + harness: _Harness, endpoint: _Endpoint, membership: str +) -> None: + if membership == "catalog": + with harness.factory.begin() as session: + session.execute(delete(TrialApp).where(TrialApp.app_id == harness.target.id)) + harness.catalog.ids.add(harness.target.id) + + response = harness.get(endpoint, query=harness.dataset_query() if endpoint == "datasets" else "") + + assert response.status_code == 200, response.get_json() + assert harness.catalog.calls == ([] if membership == "trial" else [harness.target.id]) + with harness.factory() as session: + assert session.scalar(select(AccountTrialAppRecord.count)) == 100 + + +@pytest.mark.parametrize("endpoint", ["parameters", "site", "datasets"]) +@pytest.mark.parametrize("unavailable", ["unlisted", "disabled", "missing"]) +def test_preview_rejects_unavailable_apps_before_loading_content( + harness: _Harness, endpoint: _Endpoint, unavailable: str +) -> None: + with harness.factory.begin() as session: + if unavailable == "unlisted": + session.execute(delete(TrialApp).where(TrialApp.app_id == harness.target.id)) + elif unavailable == "disabled": + session.execute( + text("UPDATE apps SET status = 'disabled' WHERE id = :app_id"), {"app_id": harness.target.id} + ) + else: + session.execute(delete(App).where(App.id == harness.target.id)) + + _assert_error(harness.get(endpoint, query=harness.dataset_query()), 404, "app_not_found") + assert harness.signed_icons == [] + + +def test_parameters_preserve_complete_shape_and_zero_false_empty_values(harness: _Harness) -> None: + response = harness.get("parameters") + assert response.status_code == 200, response.get_json() + assert response.get_json() == { + "opening_statement": "", + "suggested_questions": [], + "suggested_questions_after_answer": {"enabled": False}, + "speech_to_text": {"enabled": False}, + "text_to_speech": {"enabled": False, "voice": "", "autoPlay": "disabled"}, + "retriever_resource": {"enabled": False}, + "annotation_reply": {"enabled": False}, + "more_like_this": {"enabled": False}, + "user_input_form": _INPUT_FORM, + "sensitive_word_avoidance": {"enabled": False, "type": "", "configs": []}, + "file_upload": {"enabled": False, "number_limits": 0}, + "system_parameters": { + "image_file_size_limit": 1, + "video_file_size_limit": 2, + "audio_file_size_limit": 3, + "file_size_limit": 4, + "workflow_file_upload_limit": 5, + }, + } + + +def test_missing_parameter_configuration_is_app_unavailable(harness: _Harness) -> None: + with harness.factory.begin() as session: + session.execute(delete(AppModelConfig).where(AppModelConfig.id == harness.config.id)) + _assert_error(harness.get("parameters"), 400, "app_unavailable") + + +def test_site_preserves_complete_fields_and_signs_after_session_closes(harness: _Harness) -> None: + response = harness.get("site") + assert response.status_code == 200, response.get_json() + assert harness.target.enable_site is False + assert response.get_json() == { + "title": "Preview site", + "chat_color_theme": "", + "chat_color_theme_inverted": False, + "icon_type": "image", + "icon": harness.site.icon, + "icon_background": "", + "description": None, + "copyright": "", + "privacy_policy": None, + "input_placeholder": "", + "custom_disclaimer": "", + "default_language": "en-US", + "show_workflow_steps": False, + "use_icon_as_answer_icon": False, + "icon_url": f"https://files.example/{harness.site.icon}?sign=preview", + } + assert harness.signed_icons == [harness.site.icon] + + +@pytest.mark.parametrize("icon_type", [IconType.EMOJI, None]) +def test_non_image_site_icon_has_no_signed_url(harness: _Harness, icon_type: IconType | None) -> None: + with harness.factory.begin() as session: + session.execute(update(Site).where(Site.id == harness.site.id).values(icon_type=icon_type, icon="")) + response = harness.get("site") + assert response.status_code == 200 + assert response.get_json()["icon_url"] is None + assert harness.signed_icons == [] + + +@pytest.mark.parametrize( + ("missing", "code"), + [("site", "app_site_unavailable"), ("owner", "app_owner_unavailable"), ("archived_owner", "app_owner_unavailable")], +) +def test_site_errors_distinguish_site_and_owner_unavailability(harness: _Harness, missing: str, code: str) -> None: + with harness.factory.begin() as session: + if missing == "site": + session.execute(delete(Site).where(Site.app_id == harness.target.id)) + elif missing == "owner": + session.execute(delete(Tenant).where(Tenant.id == harness.owner.id)) + else: + session.execute(update(Tenant).where(Tenant.id == harness.owner.id).values(status=TenantStatus.ARCHIVE)) + _assert_error(harness.get("site"), 403, code) + assert harness.signed_icons == [] + + +def test_dataset_ids_filter_by_actual_app_owner_without_requiring_binding(harness: _Harness) -> None: + with harness.factory.begin() as session: + other_owner = Dataset( + tenant_id=harness.listing.tenant_id, name="Listing tenant dataset", created_by=harness.creator_id + ) + unrequested = Dataset(tenant_id=harness.owner.id, name="Not requested", created_by=harness.creator_id) + session.add_all([other_owner, unrequested]) + response = harness.get( + "datasets", query=f"?ids={harness.dataset.id}&ids={harness.dataset.id}&ids={other_owner.id}&ids={uuid4()}" + ) + assert response.status_code == 200, response.get_json() + assert response.get_json() == { + "data": [ + { + "id": harness.dataset.id, + "name": "Owner dataset", + "description": None, + "permission": "only_me", + "data_source_type": "upload_file", + "indexing_technique": None, + "created_by": harness.creator_id, + "created_at": int(_CREATED_AT.timestamp()), + "permission_keys": [], + } + ], + "has_more": False, + "limit": 20, + "total": 1, + "page": 1, + } + + +@pytest.mark.parametrize( + ("query", "page", "limit", "has_more"), + [ + ("&page=0&limit=0", 0, 0, False), + ("&page=-2&limit=-1", -2, -1, False), + ("&page=invalid&limit=invalid", 1, 20, False), + ("&page=99&limit=1", 99, 1, True), + ], +) +def test_dataset_page_and_limit_are_metadata_without_slicing( + harness: _Harness, query: str, page: int, limit: int, has_more: bool +) -> None: + response = harness.get("datasets", query=harness.dataset_query() + query) + assert response.status_code == 200, response.get_json() + body = response.get_json() + assert [item["id"] for item in body["data"]] == [harness.dataset.id] + assert {key: body[key] for key in ("page", "limit", "total", "has_more")} == { + "page": page, + "limit": limit, + "total": 1, + "has_more": has_more, + } + + +@pytest.mark.parametrize("ids", ["missing", "empty"]) +def test_unmatched_dataset_ids_return_empty_result(harness: _Harness, ids: str) -> None: + response = harness.get("datasets", query=f"?ids={uuid4() if ids == 'missing' else ''}") + assert response.status_code == 200 + assert response.get_json() == {"data": [], "has_more": False, "limit": 20, "total": 0, "page": 1} + + +def test_dataset_list_requires_ids(harness: _Harness) -> None: + _assert_error(harness.get("datasets"), 400, "need_add_ids") diff --git a/api/tests/unit_tests/extensions/test_ext_application_services.py b/api/tests/unit_tests/extensions/test_ext_application_services.py index dc435736ec225f..0f763bb57c61b5 100644 --- a/api/tests/unit_tests/extensions/test_ext_application_services.py +++ b/api/tests/unit_tests/extensions/test_ext_application_services.py @@ -64,6 +64,7 @@ RedisOAuthAccountClaimLock, ) from services.app_generate_service import AppGenerateService +from services.app_preview_query_service import AppPreviewRef, AppPreviewUnavailableError from services.app_site_service import AppSiteService from services.app_tracing_config_gateway import OpsTraceManagerGateway from services.app_tracing_config_service import AppTracingConfigService @@ -737,6 +738,33 @@ def test_trial_generation_uses_configured_access_runtime_and_usage( assert record.count == 1 +def test_app_previews_use_the_configured_catalog_and_app_owner( + sqlite_session_factory: sessionmaker[Session], monkeypatch: pytest.MonkeyPatch +) -> None: + apply_config_overrides(monkeypatch, HOSTED_FETCH_APP_TEMPLATES_MODE="builtin") + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=DeploymentEdition.COMMUNITY, + initialization_password="", + redis=MagicMock(spec=RedisClientWrapper), + ) + app_id, tenant_id, other_id = str(uuid4()), str(uuid4()), str(uuid4()) + with sqlite_session_factory.begin() as session: + session.add_all( + [ + App(id=app_id, tenant_id=tenant_id, name="Preview", mode="chat", enable_site=False, enable_api=False), + App(id=other_id, tenant_id=tenant_id, name="Private", mode="chat", enable_site=False, enable_api=False), + ] + ) + + # Catalog-only previews must work without a Trial registration or account. + payload = json.dumps({"app_details": {app_id: {"id": app_id}}}) + with patch.object(recommended_app_catalog_gateway.Path, "read_text", return_value=payload): + assert services.app_previews.get_access(app_id=app_id) == AppPreviewRef(app_id=app_id, tenant_id=tenant_id) + with pytest.raises(AppPreviewUnavailableError, match=other_id): + services.app_previews.get_access(app_id=other_id) + + def test_build_application_services_adapts_enterprise_webapp_access_mode( sqlite_session_factory: sessionmaker[Session], ) -> None: diff --git a/api/tests/unit_tests/repositories/test_app_preview_query_repository.py b/api/tests/unit_tests/repositories/test_app_preview_query_repository.py new file mode 100644 index 00000000000000..903dd4d437bf9b --- /dev/null +++ b/api/tests/unit_tests/repositories/test_app_preview_query_repository.py @@ -0,0 +1,342 @@ +"""Preview reads enforce app ownership without adding execution or dataset ACL policy.""" + +from dataclasses import replace +from datetime import datetime +from typing import Literal +from uuid import uuid4 + +import pytest +from sqlalchemy import Engine, select, text +from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.pool import QueuePool + +from models.account import Tenant, TenantStatus +from models.dataset import Dataset +from models.enums import CustomizeTokenStrategy +from models.model import App, AppMode, IconType, Site +from repositories.app_preview_query_repository import AppPreviewQueryRepository +from services.app_definition_query_service import AppSiteConfiguration +from services.app_preview_query_service import AppPreviewDataset, AppPreviewRef + +_APP_ID = "11111111-1111-1111-1111-111111111111" +_TENANT_ID = "22222222-2222-2222-2222-222222222222" +_DECOY_APP_ID = "33333333-3333-3333-3333-333333333333" +_DECOY_TENANT_ID = "44444444-4444-4444-4444-444444444444" +_DATASET_ID = "55555555-5555-5555-5555-555555555555" +_DECOY_DATASET_ID = "66666666-6666-6666-6666-666666666666" +_CREATOR_ID = "77777777-7777-7777-7777-777777777777" +_CREATED_AT = datetime(2024, 1, 2, 3, 4, 5) + + +def _add_app(session: Session, *, app_id: str, tenant_id: str) -> None: + session.add( + App( + id=app_id, + tenant_id=tenant_id, + name="Preview", + mode=AppMode.CHAT, + enable_site=False, + enable_api=False, + is_public=False, + ) + ) + + +def _add_site(session: Session, *, app_id: str, title: str) -> None: + session.add( + Site( + app_id=app_id, + title=title, + default_language="zh-Hans", + customize_token_strategy=CustomizeTokenStrategy.UUID, + icon_type=IconType.IMAGE, + icon=_DATASET_ID, + icon_background="#ffffff", + description="Preview description", + copyright="Copyright", + privacy_policy="https://example.com/privacy", + input_placeholder="Ask a question", + custom_disclaimer="Disclaimer", + chat_color_theme="#000000", + chat_color_theme_inverted=True, + prompt_public=True, + show_workflow_steps=True, + use_icon_as_answer_icon=True, + ) + ) + + +def _add_dataset(session: Session, *, dataset_id: str, tenant_id: str) -> None: + session.add( + Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Knowledge", + description="Knowledge description", + permission="only_me", + data_source_type="upload_file", + indexing_technique="high_quality", + created_by=_CREATOR_ID, + created_at=_CREATED_AT, + ) + ) + + +@pytest.fixture +def preview_ref(sqlite_session_factory: sessionmaker[Session]) -> AppPreviewRef: + with sqlite_session_factory.begin() as session: + tenant = Tenant(name="Owner") + tenant.id = _TENANT_ID + decoy_tenant = Tenant(name="Other owner") + decoy_tenant.id = _DECOY_TENANT_ID + session.add_all([tenant, decoy_tenant]) + _add_app(session, app_id=_APP_ID, tenant_id=_TENANT_ID) + _add_app(session, app_id=_DECOY_APP_ID, tenant_id=_DECOY_TENANT_ID) + _add_site(session, app_id=_APP_ID, title="Preview Site") + _add_site(session, app_id=_DECOY_APP_ID, title="Decoy Site") + _add_dataset(session, dataset_id=_DATASET_ID, tenant_id=_TENANT_ID) + _add_dataset(session, dataset_id=_DECOY_DATASET_ID, tenant_id=_DECOY_TENANT_ID) + return AppPreviewRef(app_id=_APP_ID, tenant_id=_TENANT_ID) + + +def test_get_app_returns_owner_without_requiring_publication_or_trial_registration( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef +) -> None: + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + + assert repository.get_app(app_id=_APP_ID) == preview_ref + + +@pytest.mark.parametrize("state", ["missing", "not-normal"]) +def test_get_app_does_not_substitute_another_normal_app( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef, state: str +) -> None: + app_id = preview_ref.app_id + if state == "missing": + app_id = str(uuid4()) + else: + with sqlite_session_factory.begin() as session: + # Legacy databases can contain statuses outside the current enum. + session.execute(text("UPDATE apps SET status = 'disabled' WHERE id = :app_id"), {"app_id": app_id}) + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + + assert repository.get_app(app_id=_DECOY_APP_ID) == AppPreviewRef(app_id=_DECOY_APP_ID, tenant_id=_DECOY_TENANT_ID) + assert repository.get_app(app_id=app_id) is None + + +def test_get_site_returns_detached_configuration_for_disabled_site_app( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef +) -> None: + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + + result = repository.get_site(app=preview_ref) + + assert result is not None + assert result.owner_status == "normal" + assert result.configuration == AppSiteConfiguration( + title="Preview Site", + chat_color_theme="#000000", + chat_color_theme_inverted=True, + icon_type="image", + icon=_DATASET_ID, + icon_background="#ffffff", + description="Preview description", + copyright="Copyright", + privacy_policy="https://example.com/privacy", + input_placeholder="Ask a question", + custom_disclaimer="Disclaimer", + default_language="zh-Hans", + prompt_public=True, + show_workflow_steps=True, + use_icon_as_answer_icon=True, + ) + + +@pytest.mark.parametrize("state", ["missing-app", "wrong-tenant", "not-normal", "missing-site"]) +def test_get_site_requires_complete_admitted_app_scope( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef, state: str +) -> None: + if state == "wrong-tenant": + preview_ref = replace(preview_ref, tenant_id=_DECOY_TENANT_ID) + else: + statements = { + "missing-app": "DELETE FROM apps WHERE id = :app_id", + "not-normal": "UPDATE apps SET status = 'disabled' WHERE id = :app_id", + "missing-site": "DELETE FROM sites WHERE app_id = :app_id", + } + with sqlite_session_factory.begin() as session: + session.execute(text(statements[state]), {"app_id": preview_ref.app_id}) + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + decoy_site = repository.get_site(app=AppPreviewRef(app_id=_DECOY_APP_ID, tenant_id=_DECOY_TENANT_ID)) + + assert decoy_site is not None + assert decoy_site.configuration.title == "Decoy Site" + assert repository.get_site(app=preview_ref) is None + + +@pytest.mark.parametrize("owner_state", ["missing", "archive", "normal"]) +def test_get_site_preserves_owner_state_for_service_policy( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef, owner_state: str +) -> None: + with sqlite_session_factory.begin() as session: + owner = session.get(Tenant, preview_ref.tenant_id) + assert owner is not None + if owner_state == "missing": + session.delete(owner) + else: + owner.status = TenantStatus(owner_state) + + result = AppPreviewQueryRepository(session_factory=sqlite_session_factory).get_site(app=preview_ref) + + assert result is not None + assert result.configuration.title == "Preview Site" + assert result.owner_status == (None if owner_state == "missing" else owner_state) + + +def test_get_site_does_not_filter_legacy_empty_site_status( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef +) -> None: + with sqlite_session_factory.begin() as session: + session.execute(text("UPDATE sites SET status = '' WHERE app_id = :app_id"), {"app_id": preview_ref.app_id}) + + result = AppPreviewQueryRepository(session_factory=sqlite_session_factory).get_site(app=preview_ref) + + assert result is not None + assert result.configuration.title == "Preview Site" + + +def test_get_datasets_filters_requested_ids_and_owner_without_acl_or_app_binding( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef +) -> None: + with sqlite_session_factory.begin() as session: + _add_dataset(session, dataset_id=str(uuid4()), tenant_id=_TENANT_ID) + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + + result = repository.get_datasets(app=preview_ref, ids=[_DATASET_ID, _DECOY_DATASET_ID, _DATASET_ID, str(uuid4())]) + + assert result == ( + AppPreviewDataset( + id=_DATASET_ID, + name="Knowledge", + description="Knowledge description", + permission="only_me", + data_source_type="upload_file", + indexing_technique="high_quality", + created_by=_CREATOR_ID, + created_at=_CREATED_AT, + ), + ) + assert repository.get_datasets(app=preview_ref, ids=[]) == () + assert repository.get_datasets(app=preview_ref, ids=[str(uuid4())]) == () + + +@pytest.mark.parametrize("state", ["missing-app", "wrong-tenant", "not-normal"]) +@pytest.mark.parametrize("empty_ids", [True, False]) +def test_get_datasets_requires_available_admitted_app_even_for_empty_ids( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef, state: str, empty_ids: bool +) -> None: + if state == "wrong-tenant": + preview_ref = replace(preview_ref, tenant_id=_DECOY_TENANT_ID) + else: + statement = ( + "DELETE FROM apps WHERE id = :app_id" + if state == "missing-app" + else "UPDATE apps SET status = 'disabled' WHERE id = :app_id" + ) + with sqlite_session_factory.begin() as session: + session.execute(text(statement), {"app_id": preview_ref.app_id}) + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + decoy_datasets = repository.get_datasets( + app=AppPreviewRef(app_id=_DECOY_APP_ID, tenant_id=_DECOY_TENANT_ID), ids=[_DECOY_DATASET_ID] + ) + + assert decoy_datasets is not None + assert len(decoy_datasets) == 1 + ids: list[str] = [] if empty_ids else [_DATASET_ID, _DECOY_DATASET_ID] + assert repository.get_datasets(app=preview_ref, ids=ids) is None + + +def test_get_datasets_preserves_nullable_columns(sqlite_session_factory: sessionmaker[Session]) -> None: + with sqlite_session_factory.begin() as session: + _add_app(session, app_id=_APP_ID, tenant_id=_TENANT_ID) + session.add( + Dataset(id=_DATASET_ID, tenant_id=_TENANT_ID, name="No optional configuration", created_by=_CREATOR_ID) + ) + + result = AppPreviewQueryRepository(session_factory=sqlite_session_factory).get_datasets( + app=AppPreviewRef(app_id=_APP_ID, tenant_id=_TENANT_ID), ids=[_DATASET_ID] + ) + + assert result is not None + assert len(result) == 1 + assert result[0].description is None + assert result[0].data_source_type is None + assert result[0].indexing_technique is None + + +def test_get_datasets_returns_every_requested_match_without_a_page_limit( + sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef +) -> None: + ids = [str(uuid4()) for _ in range(25)] + with sqlite_session_factory.begin() as session: + for dataset_id in ids: + _add_dataset(session, dataset_id=dataset_id, tenant_id=preview_ref.tenant_id) + + result = AppPreviewQueryRepository(session_factory=sqlite_session_factory).get_datasets(app=preview_ref, ids=ids) + + assert result is not None + assert {dataset.id for dataset in result} == set(ids) + assert len(result) == 25 + + +@pytest.mark.parametrize("method", ["app", "site", "datasets"]) +@pytest.mark.parametrize("exists", [True, False]) +def test_preview_queries_release_connections_on_found_and_missing_results( + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], + preview_ref: AppPreviewRef, + method: Literal["app", "site", "datasets"], + exists: bool, +) -> None: + if not exists: + preview_ref = replace(preview_ref, app_id=str(uuid4())) + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + assert isinstance(sqlite_engine.pool, QueuePool) + assert sqlite_engine.pool.checkedout() == 0 + + result: object + match method: + case "app": + result = repository.get_app(app_id=preview_ref.app_id) + case "site": + result = repository.get_site(app=preview_ref) + case "datasets": + result = repository.get_datasets(app=preview_ref, ids=[_DATASET_ID]) + + assert (result is not None) == exists + assert sqlite_engine.pool.checkedout() == 0 + + +def test_preview_queries_leave_caller_transaction_open_and_uncommitted( + sqlite_engine: Engine, sqlite_session_factory: sessionmaker[Session], preview_ref: AppPreviewRef +) -> None: + repository = AppPreviewQueryRepository(session_factory=sqlite_session_factory) + assert isinstance(sqlite_engine.pool, QueuePool) + with sqlite_session_factory() as caller: + app = caller.get(App, _APP_ID) + assert app is not None + app.name = "Pending caller change" + assert sqlite_engine.pool.checkedout() == 1 + + assert repository.get_app(app_id=_APP_ID) == preview_ref + assert repository.get_site(app=preview_ref) is not None + assert repository.get_datasets(app=preview_ref, ids=[_DATASET_ID]) is not None + + assert caller.in_transaction() + assert app in caller.dirty + assert sqlite_engine.pool.checkedout() == 1 + caller.rollback() + + with sqlite_session_factory() as session: + assert session.scalar(select(App.name).where(App.id == _APP_ID)) == "Preview" + assert sqlite_engine.pool.checkedout() == 0