From 08008f6305bc0b253bd3733d4901637ed3376172 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Tue, 11 Aug 2026 17:29:40 +0900 Subject: [PATCH] test: migrate controller sessions to SQLite (#40083) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../commands/test_fix_app_site_missing.py | 155 +++++++++----- .../controllers/common/test_app_access.py | 27 +-- .../console/datasets/test_external.py | 65 +++--- .../console/datasets/test_wraps.py | 7 +- .../console/explore/test_banner.py | 189 +++++++++++------- .../console/explore/test_completion.py | 7 +- .../controllers/console/explore/test_trial.py | 54 +++-- .../controllers/console/test_apikey.py | 117 +++++------ .../console/workspace/test_workspace.py | 5 +- .../controllers/inner_api/app/test_dsl.py | 158 ++++++++++----- .../controllers/inner_api/test_agent_files.py | 80 +------- .../controllers/openapi/auth/test_prepare.py | 5 +- .../openapi/test_app_describe_builder.py | 22 +- .../controllers/openapi/test_app_payloads.py | 5 +- .../test_apps_permitted_external_query.py | 5 +- .../controllers/openapi/test_input_schema.py | 41 ++-- .../controllers/service_api/app/test_audio.py | 5 +- .../service_api/app/test_conversation.py | 2 - .../service_api/app/test_hitl_service_api.py | 6 +- .../test_rag_pipeline_workflow.py | 3 - .../service_api/dataset/test_document.py | 114 +++++++---- .../service_api/dataset/test_metadata.py | 57 +++--- .../controllers/web/test_completion.py | 4 +- 23 files changed, 637 insertions(+), 496 deletions(-) diff --git a/api/tests/unit_tests/commands/test_fix_app_site_missing.py b/api/tests/unit_tests/commands/test_fix_app_site_missing.py index a7b05e3bbb1..c69f4eb6463 100644 --- a/api/tests/unit_tests/commands/test_fix_app_site_missing.py +++ b/api/tests/unit_tests/commands/test_fix_app_site_missing.py @@ -1,77 +1,132 @@ +import uuid from types import SimpleNamespace from unittest.mock import MagicMock import pytest -from sqlalchemy.orm import Session +from sqlalchemy import event +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker from commands import system as system_commands +from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole +from models.enums import CustomizeTokenStrategy +from models.model import App, AppMode, IconType, Site -def test_fix_app_site_missing_passes_loaded_session_to_signal(monkeypatch: pytest.MonkeyPatch) -> None: - account = object() - tenant = MagicMock() - tenant.get_accounts.return_value = [account] - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") +def _persist_missing_site_owner(session: Session) -> tuple[Account, App]: + """Persist an app without a Site and its complete tenant owner chain.""" + tenant = Tenant(name="Command workspace") + account = Account(name="Owner", email=f"owner-{uuid.uuid4()}@example.com") + membership = TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + current=True, + role=TenantAccountRole.OWNER, + ) + app = App( + id=str(uuid.uuid4()), + tenant_id=tenant.id, + name="Missing Site App", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#FFFFFF", + enable_site=True, + enable_api=False, + created_by=account.id, + ) + session.add_all([tenant, account, membership, app]) + session.commit() + return account, app - session = Session() + +def _site_for(app: App) -> Site: + return Site( + app_id=app.id, + title=app.name, + default_language="en-US", + customize_token_strategy=CustomizeTokenStrategy.UUID, + code=f"site-{app.id}", + ) + + +def _bind_command_database( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], +) -> tuple[scoped_session[Session], Session]: + """Expose a callable real scoped session through the Flask-SQLAlchemy shape.""" + command_sessions = scoped_session(sqlite_session_factory) + command_session = command_sessions() + monkeypatch.setattr( + system_commands, + "db", + SimpleNamespace(engine=sqlite_engine, session=command_sessions), + ) + return command_sessions, command_session + + +def test_fix_app_site_missing_passes_loaded_session_to_signal( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + account, app = _persist_missing_site_owner(sqlite_session) + command_sessions, command_session = _bind_command_database(monkeypatch, sqlite_engine, sqlite_session_factory) phase_events: list[str] = [] - scalar = MagicMock(return_value=app) - get = MagicMock(return_value=tenant) - commit = MagicMock(side_effect=lambda: phase_events.append("commit")) - monkeypatch.setattr(session, "scalar", scalar) - monkeypatch.setattr(session, "get", get) - monkeypatch.setattr(session, "commit", commit) + event.listen(command_session, "after_commit", lambda _session: phase_events.append("commit")) - scoped_session = MagicMock(return_value=session) - scoped_session.scalar.return_value = app + def create_site(sender: App, *, account: Account, session: Session) -> None: + phase_events.append("signal") + assert sender.id == app.id + assert account.id == account_id + assert session is command_session + session.add(_site_for(sender)) - connection = MagicMock() - connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []] - engine = MagicMock() - engine.begin.return_value.__enter__.return_value = connection - - monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=scoped_session)) - send = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal")) + account_id = account.id + send = MagicMock(side_effect=create_site) monkeypatch.setattr(system_commands.app_was_created, "send", send) - system_commands.fix_app_site_missing.callback() + try: + system_commands.fix_app_site_missing.callback() + finally: + command_sessions.remove() - scoped_session.assert_called_once_with() - scalar.assert_called_once() - get.assert_called_once_with(system_commands.Tenant, app.tenant_id) - tenant.get_accounts.assert_called_once_with(session=session) - send.assert_called_once_with(app, account=account, session=session) - commit.assert_called_once_with() + send.assert_called_once() assert phase_events == ["signal", "commit"] - assert isinstance(send.call_args.kwargs["session"], Session) + sqlite_session.expire_all() + persisted_site = sqlite_session.query(Site).filter_by(app_id=app.id).one() + assert persisted_site.title == app.name -def test_fix_app_site_missing_rolls_back_when_signal_fails(monkeypatch: pytest.MonkeyPatch) -> None: - account = object() - tenant = MagicMock() - tenant.get_accounts.return_value = [account] - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") - session = MagicMock() +def test_fix_app_site_missing_rolls_back_when_signal_fails( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + _account, app = _persist_missing_site_owner(sqlite_session) + command_sessions, command_session = _bind_command_database(monkeypatch, sqlite_engine, sqlite_session_factory) phase_events: list[str] = [] - session.scalar.return_value = app - session.get.return_value = tenant - session.rollback.side_effect = lambda: phase_events.append("rollback") + event.listen(command_session, "after_rollback", lambda _session: phase_events.append("rollback")) - connection = MagicMock() - connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []] - engine = MagicMock() - engine.begin.return_value.__enter__.return_value = connection - - monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=MagicMock(return_value=session))) - - def fail_signal(*_args, **_kwargs) -> None: + def fail_signal(sender: App, **_kwargs: object) -> None: phase_events.append("signal") + # Ensure the command's next raw scan terminates while its own transaction + # still exercises the rollback path. + with sqlite_session_factory() as observer: + observer.add(_site_for(sender)) + observer.commit() raise RuntimeError("failed") monkeypatch.setattr(system_commands.app_was_created, "send", MagicMock(side_effect=fail_signal)) - system_commands.fix_app_site_missing.callback() + try: + system_commands.fix_app_site_missing.callback() + finally: + command_sessions.remove() - session.rollback.assert_called_once_with() - session.commit.assert_not_called() assert phase_events == ["signal", "rollback"] + sqlite_session.expire_all() + assert sqlite_session.get(App, app.id) is not None diff --git a/api/tests/unit_tests/controllers/common/test_app_access.py b/api/tests/unit_tests/controllers/common/test_app_access.py index df5debb24d6..2a3212b5940 100644 --- a/api/tests/unit_tests/controllers/common/test_app_access.py +++ b/api/tests/unit_tests/controllers/common/test_app_access.py @@ -2,9 +2,8 @@ from __future__ import annotations -from unittest.mock import MagicMock - import pytest +from sqlalchemy.orm import Session from controllers.common.app_access import ( APP_LIST_PERMISSION_KEYS, @@ -106,28 +105,32 @@ class TestResolveAppAccessFilter: lambda tenant_id, account_id: whitelist, ) - def test_default_preview_is_unrestricted(self, monkeypatch: pytest.MonkeyPatch): + def test_default_preview_is_unrestricted(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=True)) permissions = _permissions(app_default_keys=["app.preview"]) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) assert flt.accessible_app_ids is None assert flt.can_manage_own_apps is False - def test_default_preview_overrides_whitelist_restriction(self, monkeypatch: pytest.MonkeyPatch): + def test_default_preview_overrides_whitelist_restriction( + self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=["app-9"])) permissions = _permissions( workspace_keys=["app.full_access", "app.create_and_management"], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) # Workspace-level preview grant defeats the whitelist restriction. assert flt.accessible_app_ids is None assert flt.can_manage_own_apps is True - def test_override_apps_collected_without_default_preview(self, monkeypatch: pytest.MonkeyPatch): + def test_override_apps_collected_without_default_preview( + self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=True)) permissions = _permissions( app_overrides=[ @@ -136,23 +139,23 @@ class TestResolveAppAccessFilter: ], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) assert flt.accessible_app_ids == {"app-1"} - def test_whitelist_union_with_override_apps(self, monkeypatch: pytest.MonkeyPatch): + def test_whitelist_union_with_override_apps(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=["app-5"])) permissions = _permissions( app_overrides=[ResourcePermissionKeys(resource_id="app-1", permission_keys=["app.acl.preview"])], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) assert flt.accessible_app_ids == {"app-1", "app-5"} - def test_fetches_permissions_when_not_supplied(self, monkeypatch: pytest.MonkeyPatch): + def test_fetches_permissions_when_not_supplied(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=[])) - session = MagicMock() + session = unbound_session captured: dict[str, object] = {} def get_permissions(tenant_id: str, account_id: str, *, session: object): diff --git a/api/tests/unit_tests/controllers/console/datasets/test_external.py b/api/tests/unit_tests/controllers/console/datasets/test_external.py index 305a9adbacc..c08b68c0fbc 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_external.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_external.py @@ -5,6 +5,7 @@ from unittest.mock import ANY, MagicMock, PropertyMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -162,13 +163,21 @@ def _dataset_detail_object() -> SimpleNamespace: ) -class TestExternalApiTemplateListApi: +class _UsesSQLiteSession: + session: Session + + @pytest.fixture(autouse=True) + def _inject_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + + +class TestExternalApiTemplateListApi(_UsesSQLiteSession): def test_get_success(self, app: Flask): api = ExternalApiTemplateListApi() method = inspect.unwrap(api.get) api_item = _external_api_object("api-1") - session = MagicMock() + session = self.session with ( app.test_request_context("/?page=2&limit=1&keyword=vector"), @@ -207,7 +216,7 @@ class TestExternalApiTemplateListApi: }, } created = _external_api_object("api-created") - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), @@ -247,7 +256,7 @@ class TestExternalApiTemplateListApi: patch.object(ExternalDatasetService, "validate_api_list"), ): with pytest.raises(Forbidden): - method(api, ExternalKnowledgeApiPayload.model_validate(payload), MagicMock(), "tenant-1", current_user) + method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user) def test_post_duplicate_name(self, app: Flask, current_user: Account): api = ExternalApiTemplateListApi() @@ -266,15 +275,15 @@ class TestExternalApiTemplateListApi: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, ExternalKnowledgeApiPayload.model_validate(payload), MagicMock(), "tenant-1", current_user) + method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user) -class TestExternalApiTemplateApi: +class TestExternalApiTemplateApi(_UsesSQLiteSession): def test_get_success_returns_template_contract(self, app: Flask): api = ExternalApiTemplateApi() method = inspect.unwrap(api.get) template = _external_api_object("api-detail") - session = MagicMock() + session = self.session with ( app.test_request_context("/"), @@ -306,7 +315,7 @@ class TestExternalApiTemplateApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", "api-id") + method(api, self.session, "tenant-1", "api-id") def test_patch_success_uses_validated_payload_and_returns_template(self, app: Flask, current_user: Account): api = ExternalApiTemplateApi() @@ -321,7 +330,7 @@ class TestExternalApiTemplateApi: }, } updated = _external_api_object("api-updated") - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), @@ -362,15 +371,15 @@ class TestExternalApiTemplateApi: with app.test_request_context("/"): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", current_user, "api-id") + method(api, self.session, "tenant-1", current_user, "api-id") -class TestExternalApiUseCheckApi: +class TestExternalApiUseCheckApi(_UsesSQLiteSession): def test_get_scopes_usage_check_to_current_tenant(self, app: Flask): api = ExternalApiUseCheckApi() method = inspect.unwrap(api.get) - session = MagicMock() + session = self.session with ( app.test_request_context("/"), @@ -387,7 +396,7 @@ class TestExternalApiUseCheckApi: mock_use_check.assert_called_once_with("api-id", "tenant-1", session=ANY) -class TestExternalDatasetCreateApi: +class TestExternalDatasetCreateApi(_UsesSQLiteSession): def test_create_success(self, app: Flask, current_user: Account): api = ExternalDatasetCreateApi() method = inspect.unwrap(api.post) @@ -450,10 +459,12 @@ class TestExternalDatasetCreateApi: patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), ): with pytest.raises(Forbidden): - method(api, ExternalDatasetCreatePayload.model_validate(payload), MagicMock(), "tenant-1", current_user) + method( + api, ExternalDatasetCreatePayload.model_validate(payload), self.session, "tenant-1", current_user + ) -class TestExternalKnowledgeHitTestingApi: +class TestExternalKnowledgeHitTestingApi(_UsesSQLiteSession): def test_hit_testing_dataset_not_found(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() method = inspect.unwrap(api.post) @@ -467,7 +478,7 @@ class TestExternalKnowledgeHitTestingApi: ), ): with pytest.raises(NotFound): - method(api, ExternalHitTestingPayload(query="test"), MagicMock(), current_user, "dataset-id") + method(api, ExternalHitTestingPayload(query="test"), self.session, current_user, "dataset-id") def test_hit_testing_success(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() @@ -498,7 +509,7 @@ class TestExternalKnowledgeHitTestingApi: } ], } - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), @@ -555,8 +566,6 @@ class TestBedrockRetrievalApi: ] } - session = MagicMock() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), @@ -576,7 +585,7 @@ class TestBedrockRetrievalApi: assert knowledge_id == "knowledge-base-1" -class TestExternalApiTemplateListApiAdvanced: +class TestExternalApiTemplateListApiAdvanced(_UsesSQLiteSession): def test_post_duplicate_name_error(self, app: Flask, current_user: Account): api = ExternalApiTemplateListApi() method = inspect.unwrap(api.post) @@ -593,7 +602,7 @@ class TestExternalApiTemplateListApiAdvanced: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, ExternalKnowledgeApiPayload.model_validate(payload), MagicMock(), "tenant-1", current_user) + method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user) def test_get_with_pagination(self, app: Flask): api = ExternalApiTemplateListApi() @@ -608,7 +617,7 @@ class TestExternalApiTemplateListApiAdvanced: return_value=(templates, 25), ) as get_external_knowledge_apis, ): - resp, status = method(api, ExternalApiTemplateListQuery(page=2, limit=3), MagicMock(), "tenant-1") + resp, status = method(api, ExternalApiTemplateListQuery(page=2, limit=3), self.session, "tenant-1") assert status == 200 assert resp == { @@ -621,7 +630,7 @@ class TestExternalApiTemplateListApiAdvanced: get_external_knowledge_apis.assert_called_once_with(2, 3, "tenant-1", None, session=ANY) -class TestExternalDatasetCreateApiAdvanced: +class TestExternalDatasetCreateApiAdvanced(_UsesSQLiteSession): def test_create_forbidden(self, app: Flask, current_user: Account): """Test creating external dataset without permission""" api = ExternalDatasetCreateApi() @@ -638,10 +647,12 @@ class TestExternalDatasetCreateApiAdvanced: with app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(Forbidden): - method(api, ExternalDatasetCreatePayload.model_validate(payload), MagicMock(), "tenant-1", current_user) + method( + api, ExternalDatasetCreatePayload.model_validate(payload), self.session, "tenant-1", current_user + ) -class TestExternalKnowledgeHitTestingApiAdvanced: +class TestExternalKnowledgeHitTestingApiAdvanced(_UsesSQLiteSession): def test_hit_testing_dataset_not_found(self, app: Flask, current_user: Account): """Test hit testing on non-existent dataset""" api = ExternalKnowledgeHitTestingApi() @@ -661,7 +672,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced: ), ): with pytest.raises(NotFound): - method(api, ExternalHitTestingPayload.model_validate(payload), MagicMock(), current_user, "ds-1") + method(api, ExternalHitTestingPayload.model_validate(payload), self.session, current_user, "ds-1") def test_hit_testing_with_custom_retrieval_model(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() @@ -673,7 +684,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced: "external_retrieval_model": {"type": "bm25"}, "metadata_filtering_conditions": {"status": "active"}, } - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), diff --git a/api/tests/unit_tests/controllers/console/datasets/test_wraps.py b/api/tests/unit_tests/controllers/console/datasets/test_wraps.py index ab7a7978991..0469aebda7a 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_wraps.py @@ -65,9 +65,8 @@ class TestGetRagPipeline: assert result is pipeline get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session_factory.return_value) - def test_load_rag_pipeline_uses_provided_session(self, mocker: MockerFixture): + def test_load_rag_pipeline_uses_provided_session(self, mocker: MockerFixture, sqlite_session: Session): pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") - session = Mock(spec=Session) mocker.patch( "controllers.console.datasets.wraps.current_account_with_tenant", @@ -78,10 +77,10 @@ class TestGetRagPipeline: return_value=pipeline, ) - result = load_rag_pipeline(session, "pipeline-1") + result = load_rag_pipeline(sqlite_session, "pipeline-1") assert result is pipeline - get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session) + get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=sqlite_session) def test_pipeline_id_removed_from_kwargs(self, mocker: MockerFixture): pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") diff --git a/api/tests/unit_tests/controllers/console/explore/test_banner.py b/api/tests/unit_tests/controllers/console/explore/test_banner.py index 594be28f0df..0260cffde57 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_banner.py +++ b/api/tests/unit_tests/controllers/console/explore/test_banner.py @@ -1,42 +1,31 @@ -from collections.abc import Iterator from datetime import datetime -from types import SimpleNamespace -from unittest.mock import MagicMock, call +from typing import NamedTuple from uuid import uuid4 import pytest from flask import Flask from pydantic import ValidationError -from sqlalchemy.engine import Engine -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker import controllers.console.explore.banner as banner_module -from models.base import TypeBase from models.enums import BannerStatus from models.model import ExporleBanner from repositories.explore_banner_query_repository import ExploreBannerQueryRepository from services.explore_banner_query_service import ExploreBannerQueryService, ExploreBannerRecord -@pytest.fixture -def banner_session(sqlite_engine: Engine) -> Iterator[scoped_session[Session]]: - """Create the banner table without its PostgreSQL-cast server defaults.""" - table = TypeBase.metadata.tables[ExporleBanner.__tablename__] - status_default = table.c.status.server_default - language_default = table.c.language.server_default - table.c.status.server_default = None - table.c.language.server_default = None - try: - TypeBase.metadata.create_all(sqlite_engine, tables=[table]) - finally: - table.c.status.server_default = status_default - table.c.language.server_default = language_default +class FakeExploreBannerQuery: + def __init__(self, responses: dict[str, tuple[ExploreBannerRecord, ...]] | None = None) -> None: + self.responses = responses or {} + self.requested_languages: list[str] = [] - session = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False)) - try: - yield session - finally: - session.remove() + def list_enabled(self, language: str) -> tuple[ExploreBannerRecord, ...]: + self.requested_languages.append(language) + return self.responses.get(language, ()) + + +class _ApplicationServicesStub(NamedTuple): + explore_banner_queries: ExploreBannerQueryService def _content( @@ -90,46 +79,59 @@ def _record( ) +def _use_sqlite_banner_service( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: + service = ExploreBannerQueryService( + banners=ExploreBannerQueryRepository(sqlite_session_factory), + is_enabled=lambda: True, + ) + monkeypatch.setattr( + banner_module, + "application_services", + lambda: _ApplicationServicesStub(explore_banner_queries=service), + ) + + class TestExploreBannerQueryService: def test_returns_empty_without_querying_when_disabled(self) -> None: - banners = MagicMock() + banners = FakeExploreBannerQuery() service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: False) assert service.list_for_language("fr-FR") == () - banners.list_enabled.assert_not_called() + assert banners.requested_languages == [] def test_returns_requested_language(self) -> None: record = _record() - banners = MagicMock() - banners.list_enabled.return_value = (record,) + banners = FakeExploreBannerQuery({"fr-FR": (record,)}) service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True) assert service.list_for_language("fr-FR") == (record,) - banners.list_enabled.assert_called_once_with("fr-FR") + assert banners.requested_languages == ["fr-FR"] def test_falls_back_to_en_us(self) -> None: record = _record(title="fallback") - banners = MagicMock() - banners.list_enabled.side_effect = [(), (record,)] + banners = FakeExploreBannerQuery({"en-US": (record,)}) service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True) assert service.list_for_language("es-ES") == (record,) - assert banners.list_enabled.call_args_list == [ - call("es-ES"), - call("en-US"), - ] + assert banners.requested_languages == ["es-ES", "en-US"] def test_does_not_repeat_default_language_query(self) -> None: - banners = MagicMock() - banners.list_enabled.return_value = () + banners = FakeExploreBannerQuery() service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True) assert service.list_for_language("en-US") == () - banners.list_enabled.assert_called_once_with("en-US") + assert banners.requested_languages == ["en-US"] class TestExploreBannerQueryRepository: - def test_filters_language_and_status_and_orders_by_sort(self, banner_session: scoped_session[Session]) -> None: + def test_filters_language_and_status_and_orders_by_sort( + self, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: created_at = datetime(2024, 1, 1) second = _banner( title="second", @@ -158,10 +160,10 @@ class TestExploreBannerQueryRepository: link="https://example.com/english", created_at=created_at, ) - banner_session.add_all([second, first, disabled, english]) - banner_session.commit() + sqlite_session.add_all([second, first, disabled, english]) + sqlite_session.commit() - repository = ExploreBannerQueryRepository(banner_session.session_factory) + repository = ExploreBannerQueryRepository(sqlite_session_factory) result = repository.list_enabled("fr-FR") assert [banner.id for banner in result] == [first.id, second.id] @@ -176,21 +178,29 @@ class TestExploreBannerQueryRepository: class TestBannerApi: - def test_get_serializes_requested_language(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - queries = MagicMock() - queries.list_for_language.return_value = (_record(),) - monkeypatch.setattr( - banner_module, - "application_services", - lambda: SimpleNamespace(explore_banner_queries=queries), + def test_get_serializes_requested_language( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: + banner = _banner( + title="hello", + language="fr-FR", + link="https://example.com", + created_at=datetime(2024, 1, 1), ) + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) with app.test_request_context("/?language=fr-FR"): result = banner_module.BannerApi().get() assert result == [ { - "id": "banner-1", + "id": banner.id, "content": { "category": "Featured", "title": "hello", @@ -203,31 +213,48 @@ class TestBannerApi: "created_at": "2024-01-01T00:00:00", } ] - queries.list_for_language.assert_called_once_with("fr-FR") - def test_get_uses_default_language(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - queries = MagicMock() - queries.list_for_language.return_value = () - monkeypatch.setattr( - banner_module, - "application_services", - lambda: SimpleNamespace(explore_banner_queries=queries), + def test_get_uses_default_language( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: + banner = _banner( + title="default", + language="en-US", + link="https://example.com/default", + created_at=datetime(2024, 1, 2), ) + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) with app.test_request_context("/"): result = banner_module.BannerApi().get() - assert result == [] - queries.list_for_language.assert_called_once_with("en-US") + assert result[0]["id"] == banner.id + assert result[0]["content"]["title"] == "default" - def test_get_allows_empty_supporting_copy(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - queries = MagicMock() - queries.list_for_language.return_value = (_record(category="", description=""),) - monkeypatch.setattr( - banner_module, - "application_services", - lambda: SimpleNamespace(explore_banner_queries=queries), + def test_get_allows_empty_supporting_copy( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: + banner = _banner( + title="hello", + language="en-US", + link="https://example.com", + created_at=datetime(2024, 1, 3), ) + banner.content["category"] = "" + banner.content["description"] = "" + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) with app.test_request_context("/"): result = banner_module.BannerApi().get() @@ -239,15 +266,23 @@ class TestBannerApi: "img-src": "https://example.com/banner.png", } - def test_get_rejects_invalid_content(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - queries = MagicMock() - queries.list_for_language.return_value = (_record()._replace(content={"title": "invalid"}),) - monkeypatch.setattr( - banner_module, - "application_services", - lambda: SimpleNamespace(explore_banner_queries=queries), + def test_get_rejects_invalid_content( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: + banner = _banner( + title="invalid", + language="en-US", + link="https://example.com", + created_at=datetime(2024, 1, 4), ) + banner.content = {"title": "invalid"} + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) - with app.test_request_context("/"): - with pytest.raises(ValidationError): - banner_module.BannerApi().get() + with app.test_request_context("/"), pytest.raises(ValidationError): + banner_module.BannerApi().get() diff --git a/api/tests/unit_tests/controllers/console/explore/test_completion.py b/api/tests/unit_tests/controllers/console/explore/test_completion.py index 96237934355..9ff013e27a2 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_completion.py +++ b/api/tests/unit_tests/controllers/console/explore/test_completion.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, PropertyMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError import controllers.console.explore.completion as completion_module @@ -392,7 +393,9 @@ class TestChatApi: api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app ) - def test_invalid_conversation_id_fails_fast_as_not_found(self, app: Flask, chat_app, user) -> None: + def test_invalid_conversation_id_fails_fast_as_not_found( + self, app: Flask, chat_app, user, unbound_session: Session + ) -> None: # A nonexistent conversation_id must fail fast as 404, before the streaming # generator is created. Previously the lookup only ran inside the generator, # so an invalid id surfaced as a hang instead of a clean error. @@ -407,7 +410,7 @@ class TestChatApi: get_conversation_mock = MagicMock( side_effect=completion_module.services.errors.conversation.ConversationNotExistsError() ) - session = MagicMock() + session = unbound_session api = completion_module.ChatApi() method = unwrap(api.post) 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 9406ad68064..d8838a76251 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -131,7 +131,7 @@ 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): +def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask, unbound_session: Session): class DatasetListItem: id = "dataset-1" name = "Dataset" @@ -150,8 +150,6 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): api = module.DatasetListApi() method = unwrap(api.get) app_model = SimpleNamespace(tenant_id="tenant-1") - session = MagicMock() - with ( app.test_request_context("/?page=1&limit=20&ids=dataset-1"), patch.object( @@ -160,9 +158,9 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): return_value=([DatasetListItem()], 1), ) as get_datasets, ): - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) - get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=session) + get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=unbound_session) assert result == { "data": [ { @@ -195,8 +193,9 @@ def test_trial_app_handlers_use_explicit_read_session(api_type: type) -> None: assert tuple(signature(api_type.get).parameters)[:3] == ("self", "session", "app_model") -def test_trial_app_detail_serializes_with_explicit_session(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - session = MagicMock() +def test_trial_app_detail_serializes_with_explicit_session( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: app_model = MagicMock() response_view = MagicMock() get_app = MagicMock(return_value=app_model) @@ -208,11 +207,11 @@ def test_trial_app_detail_serializes_with_explicit_session(app: Flask, monkeypat monkeypatch.setattr(module.TrialAppDetailResponse, "model_validate", MagicMock(return_value=validated)) with app.test_request_context("/"): - result = unwrap(module.AppApi.get)(module.AppApi(), session, app_model) + result = unwrap(module.AppApi.get)(module.AppApi(), unbound_session, app_model) assert result == {"id": "app-1"} - get_app.assert_called_once_with(app_model, session=session) - build_view.assert_called_once_with(app_model, session=session) + get_app.assert_called_once_with(app_model, session=unbound_session) + build_view.assert_called_once_with(app_model, session=unbound_session) module.TrialAppDetailResponse.model_validate.assert_called_once_with(response_view, from_attributes=True) @@ -882,14 +881,14 @@ class TestTrialMessageSuggestedQuestionApi: class TestTrialAppParameterApi: - def test_app_unavailable(self) -> None: + def test_app_unavailable(self, unbound_session: Session) -> None: api = module.TrialAppParameterApi() method = unwrap(api.get) with pytest.raises(AppUnavailableError): - method(api, MagicMock(), None) + method(api, unbound_session, None) - def test_success_non_workflow(self, valid_parameters: dict[str, object]) -> None: + def test_success_non_workflow(self, valid_parameters: dict[str, object], unbound_session: Session) -> None: api = module.TrialAppParameterApi() method = unwrap(api.get) @@ -899,7 +898,6 @@ class TestTrialAppParameterApi: mode=AppMode.CHAT, app_model_config_with_session=MagicMock(return_value=app_model_config), ) - session = MagicMock() annotation_reply = {"enabled": False} with ( @@ -917,14 +915,14 @@ class TestTrialAppParameterApi: return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}), ), ): - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) assert result == {"ok": True} - app_model.app_model_config_with_session.assert_called_once_with(session=session) - load_annotation_reply.assert_called_once_with(session, "app-1") + app_model.app_model_config_with_session.assert_called_once_with(session=unbound_session) + load_annotation_reply.assert_called_once_with(unbound_session, "app-1") app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) - def test_success_workflow(self, valid_parameters: dict[str, object]) -> None: + def test_success_workflow(self, valid_parameters: dict[str, object], unbound_session: Session) -> None: api = module.TrialAppParameterApi() method = unwrap(api.get) @@ -934,8 +932,6 @@ class TestTrialAppParameterApi: mode=AppMode.WORKFLOW, workflow_with_session=MagicMock(return_value=workflow), ) - session = MagicMock() - with ( patch.object(module, "get_parameters_from_feature_dict", return_value=valid_parameters), patch.object( @@ -944,10 +940,10 @@ class TestTrialAppParameterApi: return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}), ), ): - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) assert result == {"ok": True} - app_model.workflow_with_session.assert_called_once_with(session=session) + app_model.workflow_with_session.assert_called_once_with(session=unbound_session) workflow.user_input_form.assert_called_once_with(to_old_structure=True) @@ -1446,7 +1442,7 @@ class TestTrialSitApi: class TestAppWorkflowApi: - def test_uses_injected_session(self) -> None: + def test_uses_injected_session(self, unbound_session: Session) -> None: api = module.AppWorkflowApi() method = unwrap(api.get) created_by = SimpleNamespace(id="account-1", name="Creator", email="creator@example.com") @@ -1489,9 +1485,7 @@ class TestAppWorkflowApi: workflow_id="workflow-1", workflow_with_session=MagicMock(return_value=workflow), ) - session = MagicMock() - - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) assert result == { "id": "workflow-1", @@ -1535,10 +1529,10 @@ class TestAppWorkflowApi: ], "rag_pipeline_variables": [], } - app_model.workflow_with_session.assert_called_once_with(session=session) - workflow.get_created_by_account.assert_called_once_with(session=session) - workflow.get_updated_by_account.assert_called_once_with(session=session) - workflow.get_tool_published.assert_called_once_with(session=session) + app_model.workflow_with_session.assert_called_once_with(session=unbound_session) + workflow.get_created_by_account.assert_called_once_with(session=unbound_session) + workflow.get_updated_by_account.assert_called_once_with(session=unbound_session) + workflow.get_tool_published.assert_called_once_with(session=unbound_session) class TestTrialChatAudioApiExceptionHandlers: diff --git a/api/tests/unit_tests/controllers/console/test_apikey.py b/api/tests/unit_tests/controllers/console/test_apikey.py index 36f1935caf2..a2dca678c4b 100644 --- a/api/tests/unit_tests/controllers/console/test_apikey.py +++ b/api/tests/unit_tests/controllers/console/test_apikey.py @@ -2,18 +2,19 @@ from __future__ import annotations import inspect from collections.abc import Callable -from types import SimpleNamespace from typing import cast -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest +from sqlalchemy import event, select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden from controllers.console.apikey import BaseApiKeyListResource, BaseApiKeyResource from models import Account from models.account import AccountStatus, TenantAccountRole from models.enums import ApiTokenType -from models.model import ApiToken, App, AppMode +from models.model import ApiToken, App, AppMode, IconType from services.agent.errors import AgentAccessNotReadyError @@ -45,77 +46,80 @@ def _make_account(role: TenantAccountRole) -> Account: return account -def test_list_api_keys_uses_injected_session_and_tenant_id() -> None: +def _persist_app(session: Session, *, mode: AppMode = AppMode.CHAT) -> App: + app = App( + id="app-1", + tenant_id="tenant-1", + name="API key app", + mode=mode, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#ffffff", + enable_site=False, + enable_api=True, + ) + session.add(app) + session.flush() + return app + + +def test_list_api_keys_uses_injected_session_and_tenant_id(sqlite_session: Session) -> None: resource = _make_list_resource() raw_get = cast( Callable[[BaseApiKeyListResource, object, str, str], dict[str, object]], inspect.unwrap(BaseApiKeyListResource.get), ) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() - api_key = SimpleNamespace( - id="key-1", + session = sqlite_session + _persist_app(session) + api_key = ApiToken( type=ApiTokenType.APP, token="app-token", - last_used_at=None, - created_at=None, + app_id="app-1", + tenant_id="tenant-1", ) - session.scalars.return_value.all.return_value = [api_key] + api_key.id = "key-1" + session.add(api_key) + session.commit() result = raw_get(resource, session, "app-1", "tenant-1") + data = cast(list[dict[str, object]], result["data"]) - session.execute.assert_called_once() - session.scalars.assert_called_once() - assert result == { - "data": [ - { - "id": "key-1", - "type": "app", - "token": "app-token", - "last_used_at": None, - "created_at": None, - } - ] - } + assert len(data) == 1 + assert data[0]["id"] == "key-1" + assert data[0]["token"] == "app-token" -def test_create_api_key_uses_injected_session_and_tenant_id() -> None: +def test_create_api_key_uses_injected_session_and_tenant_id(sqlite_session: Session) -> None: resource = _make_list_resource() raw_post = cast( Callable[[BaseApiKeyListResource, object, str, str], tuple[dict[str, object], int]], inspect.unwrap(BaseApiKeyListResource.post), ) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() - session.scalar.return_value = 0 - - def add_api_token(api_token: ApiToken) -> None: - api_token.id = "key-1" + session = sqlite_session + _persist_app(session) + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) with patch( "controllers.console.apikey.ApiToken.generate_api_key", return_value="app-generated-token" ) as generate_api_key: - session.add.side_effect = add_api_token - result, status = raw_post(resource, session, "app-1", "tenant-1") assert status == 201 assert result["token"] == "app-generated-token" - api_token = session.add.call_args.args[0] + api_token = session.scalar(select(ApiToken).where(ApiToken.token == "app-generated-token")) + assert api_token is not None assert api_token.app_id == "app-1" assert api_token.tenant_id == "tenant-1" assert api_token.type == ApiTokenType.APP generate_api_key.assert_called_once_with("app-", 24, session=session) - session.execute.assert_called_once() - session.scalar.assert_called_once() - session.commit.assert_called_once() + assert commits == ["commit"] -def test_create_agent_api_key_requires_published_access() -> None: +def test_create_agent_api_key_requires_published_access(sqlite_session: Session) -> None: resource = _make_list_resource() - app = App(id="app-1", tenant_id="tenant-1", mode=AppMode.AGENT) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = app + session = sqlite_session + app = _persist_app(session, mode=AppMode.AGENT) with patch( "controllers.console.apikey.AppService.ensure_agent_app_access_ready", @@ -125,18 +129,17 @@ def test_create_agent_api_key_requires_published_access() -> None: resource._create_api_key("app-1", "tenant-1", session=session) ensure_access_ready.assert_called_once_with(app, session=session) - session.scalar.assert_not_called() - session.add.assert_not_called() + assert session.scalar(select(ApiToken)) is None -def test_delete_api_key_rejects_non_admin_account() -> None: +def test_delete_api_key_rejects_non_admin_account(sqlite_session: Session) -> None: resource = _make_key_resource() raw_delete = cast( Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]], inspect.unwrap(BaseApiKeyResource.delete), ) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() + session = sqlite_session + _persist_app(session) with pytest.raises(Forbidden): raw_delete( @@ -148,20 +151,21 @@ def test_delete_api_key_rejects_non_admin_account() -> None: _make_account(TenantAccountRole.NORMAL), ) - session.execute.assert_called_once() - session.scalar.assert_not_called() - -def test_delete_api_key_uses_injected_session_user_and_tenant() -> None: +def test_delete_api_key_uses_injected_session_user_and_tenant(sqlite_session: Session) -> None: resource = _make_key_resource() raw_delete = cast( Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]], inspect.unwrap(BaseApiKeyResource.delete), ) - api_key = SimpleNamespace(token="app-token", type=ApiTokenType.APP) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() - session.scalar.return_value = api_key + session = sqlite_session + _persist_app(session) + api_key = ApiToken(type=ApiTokenType.APP, token="app-token", app_id="app-1", tenant_id="tenant-1") + api_key.id = "key-1" + session.add(api_key) + session.commit() + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) with patch("controllers.console.apikey.ApiTokenCache.delete") as delete_cache: result, status = raw_delete( @@ -174,8 +178,7 @@ def test_delete_api_key_uses_injected_session_user_and_tenant() -> None: ) delete_cache.assert_called_once_with("app-token", ApiTokenType.APP) - assert session.execute.call_count == 2 - session.scalar.assert_called_once() - session.commit.assert_called_once() + assert session.get(ApiToken, "key-1") is None + assert commits == ["commit"] assert result == "" assert status == 204 diff --git a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py index c6028dba20e..dac0b09241d 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py @@ -641,9 +641,8 @@ class TestWorkspaceInfoApi: ), ), ): - session = MagicMock() - session.get.return_value = tenant - session.commit.side_effect = lambda: events.append("commit") + session = workspace_session() + event.listen(session, "after_commit", lambda _session: events.append("commit")) result = method(api, session, "t1") assert result["result"] == "success" assert events == ["commit", "get_tenant_info"] diff --git a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py index c7f788dcb55..8caeb7174ce 100644 --- a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py +++ b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py @@ -6,22 +6,47 @@ in test_auth_wraps.py; handler tests use inspect.unwrap() to bypass them. """ import inspect -from unittest.mock import ANY, MagicMock, patch +from types import SimpleNamespace +from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest from flask import Flask from pydantic import ValidationError +from sqlalchemy import event, select +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker +from controllers.inner_api.app import dsl as dsl_module from controllers.inner_api.app.dsl import ( EnterpriseAppDSLExport, EnterpriseAppDSLImport, InnerAppDSLImportPayload, _get_active_account, ) +from models import Account, App from models.account import AccountStatus +from models.model import AppMode, IconType from services.app_dsl_service import Import, ImportStatus +def _persist_app(session: Session) -> App: + app = App( + id=str(uuid4()), + tenant_id=str(uuid4()), + name="DSL App", + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#ffffff", + enable_site=False, + enable_api=False, + ) + session.add(app) + session.commit() + return app + + class TestInnerAppDSLImportPayload: """Test InnerAppDSLImportPayload Pydantic model validation.""" @@ -61,32 +86,29 @@ class TestInnerAppDSLImportPayload: class TestGetActiveAccount: """Test the _get_active_account helper function.""" - @patch("controllers.inner_api.app.dsl.db") - def test_returns_active_account(self, mock_db): - mock_account = MagicMock() - mock_account.status = AccountStatus.ACTIVE - mock_db.session.scalar.return_value = mock_account + def test_returns_active_account(self, sqlite_session: Session): + account = Account(name="Active", email="user@example.com", status=AccountStatus.ACTIVE) + sqlite_session.add(account) + sqlite_session.commit() - result = _get_active_account("user@example.com") + with patch.object(dsl_module.db, "session", sqlite_session): + result = _get_active_account("user@example.com") - assert result is mock_account - mock_db.session.scalar.assert_called_once() + assert result is account - @patch("controllers.inner_api.app.dsl.db") - def test_returns_none_for_inactive_account(self, mock_db): - mock_account = MagicMock() - mock_account.status = AccountStatus.BANNED - mock_db.session.scalar.return_value = mock_account + def test_returns_none_for_inactive_account(self, sqlite_session: Session): + account = Account(name="Banned", email="banned@example.com", status=AccountStatus.BANNED) + sqlite_session.add(account) + sqlite_session.commit() - result = _get_active_account("banned@example.com") + with patch.object(dsl_module.db, "session", sqlite_session): + result = _get_active_account("banned@example.com") assert result is None - @patch("controllers.inner_api.app.dsl.db") - def test_returns_none_for_nonexistent_email(self, mock_db): - mock_db.session.scalar.return_value = None - - result = _get_active_account("missing@example.com") + def test_returns_none_for_nonexistent_email(self, sqlite_session: Session): + with patch.object(dsl_module.db, "session", sqlite_session): + result = _get_active_account("missing@example.com") assert result is None @@ -102,20 +124,29 @@ class TestEnterpriseAppDSLImport: return EnterpriseAppDSLImport() @pytest.fixture - def _mock_import_deps(self): - """Patch db, Session, and AppDslService for import handler tests.""" - mock_session = MagicMock() - mock_session.__enter__ = MagicMock(return_value=mock_session) - mock_session.__exit__ = MagicMock(return_value=False) + def _mock_import_deps(self, sqlite_engine: Engine): + """Bind the handler Session to SQLite and isolate the DSL service boundary.""" + self._transaction_events: list[str] = [] + + def on_commit(session: Session) -> None: + if session.get_bind() is sqlite_engine: + self._transaction_events.append("commit") + + def on_rollback(session: Session) -> None: + if session.get_bind() is sqlite_engine: + self._transaction_events.append("rollback") + + event.listen(Session, "after_commit", on_commit) + event.listen(Session, "after_rollback", on_rollback) with ( - patch("controllers.inner_api.app.dsl.db"), - patch("controllers.inner_api.app.dsl.Session", return_value=mock_session), + patch.object(dsl_module, "db", SimpleNamespace(engine=sqlite_engine)), patch("controllers.inner_api.app.dsl.AppDslService") as mock_dsl_cls, ): - self._mock_session = mock_session self._mock_dsl = MagicMock() mock_dsl_cls.return_value = self._mock_dsl yield + event.remove(Session, "after_commit", on_commit) + event.remove(Session, "after_rollback", on_rollback) def _make_import_result(self, status: ImportStatus, **kwargs) -> Import: result = Import( @@ -145,9 +176,9 @@ class TestEnterpriseAppDSLImport: body, status_code = result assert status_code == 200 assert body["status"] == "completed" - mock_account.set_tenant_id_with_session.assert_called_once_with("ws-123", session=self._mock_session) - self._mock_session.commit.assert_called_once_with() - self._mock_session.rollback.assert_not_called() + call_session = mock_account.set_tenant_id_with_session.call_args.kwargs["session"] + assert isinstance(call_session, Session) + assert self._transaction_events == ["commit"] @pytest.mark.usefixtures("_mock_import_deps") @patch("controllers.inner_api.app.dsl._get_active_account") @@ -163,13 +194,14 @@ class TestEnterpriseAppDSLImport: assert status_code == 202 assert body["status"] == "pending" - self._mock_session.commit.assert_called_once_with() - self._mock_session.rollback.assert_not_called() + assert self._transaction_events == ["commit"] @pytest.mark.usefixtures("_mock_import_deps") @patch("controllers.inner_api.app.dsl._get_active_account") def test_import_failed_returns_400(self, mock_get_account, api_instance, app: Flask): - mock_get_account.return_value = MagicMock() + mock_account = MagicMock() + mock_account.set_tenant_id_with_session.side_effect = lambda _tenant_id, *, session: session.execute(select(1)) + mock_get_account.return_value = mock_account self._mock_dsl.import_app.return_value = self._make_import_result(ImportStatus.FAILED) unwrapped = inspect.unwrap(api_instance.post) @@ -180,8 +212,7 @@ class TestEnterpriseAppDSLImport: assert status_code == 400 assert body["status"] == "failed" - self._mock_session.rollback.assert_called_once_with() - self._mock_session.commit.assert_not_called() + assert self._transaction_events == ["rollback"] @patch("controllers.inner_api.app.dsl._get_active_account") def test_import_account_not_found_returns_404(self, mock_get_account, api_instance, app: Flask): @@ -208,44 +239,65 @@ class TestEnterpriseAppDSLExport: def api_instance(self): return EnterpriseAppDSLExport() + @pytest.fixture + def scoped_db(self, sqlite_session_factory: sessionmaker[Session]): + db_session = scoped_session(sqlite_session_factory) + with patch.object(dsl_module, "db", SimpleNamespace(session=db_session)): + yield db_session + db_session.remove() + @patch("controllers.inner_api.app.dsl.AppDslService") - @patch("controllers.inner_api.app.dsl.db") - def test_export_success_returns_200(self, mock_db, mock_dsl_cls, api_instance, app: Flask): - mock_app = MagicMock() - mock_db.session.get.return_value = mock_app + def test_export_success_returns_200( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) mock_dsl_cls.export_dsl.return_value = "version: 0.6.0\nkind: app\n" unwrapped = inspect.unwrap(api_instance.get) with app.test_request_context("?include_secret=false"): - result = unwrapped(api_instance, app_id="app-123") + result = unwrapped(api_instance, app_id=app_model.id) body, status_code = result assert status_code == 200 assert body["data"] == "version: 0.6.0\nkind: app\n" - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=False) + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is False @patch("controllers.inner_api.app.dsl.AppDslService") - @patch("controllers.inner_api.app.dsl.db") - def test_export_with_secret(self, mock_db, mock_dsl_cls, api_instance, app: Flask): - mock_app = MagicMock() - mock_db.session.get.return_value = mock_app + def test_export_with_secret( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) mock_dsl_cls.export_dsl.return_value = "yaml-data" unwrapped = inspect.unwrap(api_instance.get) with app.test_request_context("?include_secret=true"): - result = unwrapped(api_instance, app_id="app-123") + result = unwrapped(api_instance, app_id=app_model.id) body, status_code = result assert status_code == 200 - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=True) - - @patch("controllers.inner_api.app.dsl.db") - def test_export_app_not_found_returns_404(self, mock_db, api_instance, app: Flask): - mock_db.session.get.return_value = None + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is True + def test_export_app_not_found_returns_404(self, api_instance, app: Flask, scoped_db): + assert scoped_db() is not None unwrapped = inspect.unwrap(api_instance.get) with app.test_request_context("?include_secret=false"): - result = unwrapped(api_instance, app_id="nonexistent") + result = unwrapped(api_instance, app_id=str(uuid4())) body, status_code = result assert status_code == 404 diff --git a/api/tests/unit_tests/controllers/inner_api/test_agent_files.py b/api/tests/unit_tests/controllers/inner_api/test_agent_files.py index a34ca510220..d6a90456129 100644 --- a/api/tests/unit_tests/controllers/inner_api/test_agent_files.py +++ b/api/tests/unit_tests/controllers/inner_api/test_agent_files.py @@ -6,10 +6,10 @@ from unittest.mock import MagicMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from controllers.inner_api.agent.files import ( AgentFileDownloadRequestApi, - AgentFileRequestHttpError, AgentFileUploadRequestApi, ) from core.workflow.file_reference import build_file_reference @@ -22,7 +22,7 @@ def _raw[R](method: Callable[..., R]) -> Callable[..., R]: return cast(Callable[..., R], inspect.unwrap(method)) -def test_upload_request_returns_origin_free_uri(app: Flask) -> None: +def test_upload_request_returns_origin_free_uri(app: Flask, unbound_session: Session) -> None: payload = { "tenant_id": "tenant-1", "user_id": "execution-user-1", @@ -32,7 +32,7 @@ def test_upload_request_returns_origin_free_uri(app: Flask) -> None: } tenant = SimpleNamespace(id="tenant-1") user = SimpleNamespace(id="canonical-end-user-1") - session = MagicMock() + session = unbound_session with app.test_request_context("/", method="POST", json=payload): with ( patch(f"{MODULE}.TenantService") as tenant_service, @@ -54,71 +54,7 @@ def test_upload_request_returns_origin_free_uri(app: Flask) -> None: ) -def test_upload_request_preserves_tenant_scoped_account_owner(app: Flask) -> None: - payload = { - "tenant_id": "tenant-1", - "user_id": "account-1", - "user_from": "account", - "filename": "report.pdf", - "mimetype": "application/pdf", - "conversation_id": "conversation-1", - } - tenant = SimpleNamespace(id="tenant-1") - session = MagicMock() - with app.test_request_context("/", method="POST", json=payload): - with ( - patch(f"{MODULE}.TenantService") as tenant_service, - patch(f"{MODULE}.get_user") as get_user, - patch(f"{MODULE}.get_signed_file_uri_for_plugin", return_value="/files/upload/for-plugin?sign=1") as sign, - ): - tenant_service.get_tenant_by_id.return_value = tenant - tenant_service.account_belongs_to_tenant.return_value = True - response = _raw(AgentFileUploadRequestApi.post)(AgentFileUploadRequestApi(), session) - - assert response == {"upload_uri": "/files/upload/for-plugin?sign=1"} - get_user.assert_not_called() - tenant_service.account_belongs_to_tenant.assert_called_once_with("account-1", "tenant-1", session=session) - sign.assert_called_once_with( - filename="report.pdf", - mimetype="application/pdf", - tenant_id="tenant-1", - user_id="account-1", - conversation_id="conversation-1", - user_from="account", - ) - - -def test_upload_request_rejects_account_outside_tenant_without_signing(app: Flask) -> None: - payload = { - "tenant_id": "tenant-1", - "user_id": "account-outside-tenant", - "user_from": "account", - "filename": "report.pdf", - "mimetype": "application/pdf", - } - tenant = SimpleNamespace(id="tenant-1") - session = MagicMock() - with app.test_request_context("/", method="POST", json=payload): - with ( - patch(f"{MODULE}.TenantService") as tenant_service, - patch(f"{MODULE}.get_user") as get_user, - patch(f"{MODULE}.get_signed_file_uri_for_plugin") as sign, - ): - tenant_service.get_tenant_by_id.return_value = tenant - tenant_service.account_belongs_to_tenant.return_value = False - with pytest.raises(AgentFileRequestHttpError) as exc_info: - _raw(AgentFileUploadRequestApi.post)(AgentFileUploadRequestApi(), session) - - assert exc_info.value.error_code == "user_not_found" - assert exc_info.value.code == 404 - tenant_service.account_belongs_to_tenant.assert_called_once_with( - "account-outside-tenant", "tenant-1", session=session - ) - get_user.assert_not_called() - sign.assert_not_called() - - -def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask) -> None: +def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask, unbound_session: Session) -> None: reference = build_file_reference(record_id="tool-file-1") payload = { "tenant_id": "tenant-1", @@ -128,7 +64,7 @@ def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask) -> Non "file": {"transfer_method": "tool_file", "reference": reference}, "for_frontend": False, } - session = MagicMock() + session = unbound_session with app.test_request_context("/", method="POST", json=payload): with ( patch(f"{MODULE}.TenantService") as tenant_service, @@ -158,7 +94,9 @@ def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask) -> Non ) -def test_download_request_binds_frontend_url(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_download_request_binds_frontend_url( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: reference = build_file_reference(record_id="tool-file-1") payload = { "tenant_id": "tenant-1", @@ -169,7 +107,7 @@ def test_download_request_binds_frontend_url(app: Flask, monkeypatch: pytest.Mon "for_frontend": True, } monkeypatch.setattr(f"{MODULE}.dify_config.FILES_URL", "https://files.example.com") - session = MagicMock() + session = unbound_session with app.test_request_context("/", method="POST", json=payload): with ( patch(f"{MODULE}.TenantService") as tenant_service, diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py b/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py index 3b714a84da1..0fc691152f2 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound, Unauthorized from controllers.openapi.auth.data import AuthData, ExternalIdentity @@ -150,10 +151,10 @@ def test_load_account_skips_when_already_set(): assert data.caller is existing_caller -def test_load_account_sets_current_tenant_when_tenant_present(): +def test_load_account_sets_current_tenant_when_tenant_present(sqlite_session: Session): account = MagicMock() tenant = MagicMock() - session = MagicMock() + session = sqlite_session data = _make_auth_data(account_id=uuid.uuid4(), tenant=tenant) with ( patch("controllers.openapi.auth.prepare.AccountService.get_account_by_id", return_value=account), diff --git a/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py b/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py index f5e9b16d29c..e852189b76b 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py @@ -1,6 +1,8 @@ from types import SimpleNamespace from unittest.mock import MagicMock +from sqlalchemy.orm import Session + from controllers.openapi._input_schema import EMPTY_INPUT_SCHEMA from controllers.openapi.apps import _EMPTY_PARAMETERS, build_app_describe_response from controllers.service_api.app.error import AppUnavailableError @@ -25,9 +27,9 @@ def _app() -> _FakeApp: ) -def test_fields_none_returns_all_blocks(monkeypatch): +def test_fields_none_returns_all_blocks(monkeypatch, unbound_session: Session): app = _app() - session = MagicMock() + session = unbound_session parameters_payload = MagicMock(return_value={"k": "v"}) input_schema = MagicMock(return_value={"s": 1}) monkeypatch.setattr("controllers.openapi.apps.parameters_payload", parameters_payload) @@ -41,8 +43,8 @@ def test_fields_none_returns_all_blocks(monkeypatch): input_schema.assert_called_once_with(app, session=session) -def test_fields_subset_limits_blocks(monkeypatch): - session = MagicMock() +def test_fields_subset_limits_blocks(monkeypatch, unbound_session: Session): + session = unbound_session monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"})) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1})) resp = build_app_describe_response(_app(), ["info"], session=session) @@ -51,8 +53,8 @@ def test_fields_subset_limits_blocks(monkeypatch): assert resp.input_schema is None -def test_info_omits_author_and_tags(monkeypatch): - session = MagicMock() +def test_info_omits_author_and_tags(monkeypatch, unbound_session: Session): + session = unbound_session monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={})) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={})) resp = build_app_describe_response(_app(), ["info"], session=session) @@ -62,21 +64,21 @@ def test_info_omits_author_and_tags(monkeypatch): assert not hasattr(resp.info, "tags") -def test_parameters_fallback_on_app_unavailable(monkeypatch): +def test_parameters_fallback_on_app_unavailable(monkeypatch, unbound_session: Session): def _raise(app, *, session): raise AppUnavailableError() monkeypatch.setattr("controllers.openapi.apps.parameters_payload", _raise) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1})) - resp = build_app_describe_response(_app(), ["parameters"], session=MagicMock()) + resp = build_app_describe_response(_app(), ["parameters"], session=unbound_session) assert resp.parameters == dict(_EMPTY_PARAMETERS) -def test_input_schema_fallback_on_app_unavailable(monkeypatch): +def test_input_schema_fallback_on_app_unavailable(monkeypatch, unbound_session: Session): def _raise(app, *, session): raise AppUnavailableError() monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"})) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", _raise) - resp = build_app_describe_response(_app(), ["input_schema"], session=MagicMock()) + resp = build_app_describe_response(_app(), ["input_schema"], session=unbound_session) assert resp.input_schema == dict(EMPTY_INPUT_SCHEMA) diff --git a/api/tests/unit_tests/controllers/openapi/test_app_payloads.py b/api/tests/unit_tests/controllers/openapi/test_app_payloads.py index 2e9e7bc06a8..083a7bd980b 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_payloads.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_payloads.py @@ -8,6 +8,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from controllers.openapi.apps import ( # pyright: ignore[reportPrivateUsage] _EMPTY_PARAMETERS, @@ -35,10 +36,10 @@ def _fake_app(**overrides): return SimpleNamespace(**base) -def test_parameters_payload_raises_app_unavailable_when_no_config(): +def test_parameters_payload_raises_app_unavailable_when_no_config(unbound_session: Session): app = _fake_app(mode="chat") app.app_model_config_with_session = MagicMock(return_value=None) - session = MagicMock() + session = unbound_session with pytest.raises(AppUnavailableError): parameters_payload(app, session=session) diff --git a/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py b/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py index 49f7bea5cd3..038a17ad1df 100644 --- a/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py +++ b/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py @@ -14,6 +14,7 @@ from unittest.mock import MagicMock, patch import pytest from pydantic import ValidationError +from sqlalchemy.orm import Session from controllers.openapi.apps_permitted_external import ( PermittedExternalAppDescribeApi, @@ -69,10 +70,10 @@ def test_query_accepts_valid_mode(): assert q.mode.value == "chat" -def test_describe_forwards_request_session_to_response_builder(): +def test_describe_forwards_request_session_to_response_builder(unbound_session: Session): api = PermittedExternalAppDescribeApi() method = inspect.unwrap(api.get) - session = MagicMock() + session = unbound_session app = MagicMock() auth_data = SimpleNamespace(app=app) query = SimpleNamespace(fields={"info"}) diff --git a/api/tests/unit_tests/controllers/openapi/test_input_schema.py b/api/tests/unit_tests/controllers/openapi/test_input_schema.py index 133072ad33e..bb042cad089 100644 --- a/api/tests/unit_tests/controllers/openapi/test_input_schema.py +++ b/api/tests/unit_tests/controllers/openapi/test_input_schema.py @@ -3,6 +3,7 @@ from __future__ import annotations import pytest +from sqlalchemy.orm import Session from controllers.openapi._input_schema import _form_to_jsonschema @@ -97,6 +98,7 @@ from models.model import AppMode def _stub_app(mode: AppMode, *, form: list[dict] | None = None, has_workflow: bool | None = None): """Returns a MagicMock whose explicit config getters are wired up.""" app = MagicMock() + app.id = "00000000-0000-0000-0000-000000000001" app.mode = mode if mode in (AppMode.WORKFLOW, AppMode.ADVANCED_CHAT): if has_workflow is False: @@ -111,20 +113,15 @@ def _stub_app(mode: AppMode, *, form: list[dict] | None = None, has_workflow: bo app.app_model_config_with_session.return_value = None else: app_model_config = MagicMock() + app_model_config.app_id = app.id app_model_config.to_dict.return_value = {"user_input_form": form or []} app.app_model_config_with_session.return_value = app_model_config return app -def _session() -> MagicMock: - session = MagicMock() - session.scalar.return_value = None - return session - - -def test_chat_mode_includes_query() -> None: +def test_chat_mode_includes_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.CHAT, form=[{"text-input": {"variable": "x", "label": "X", "required": True}}]) - session = _session() + session = sqlite_session schema = build_input_schema(app, session=session) assert schema["$schema"] == "https://json-schema.org/draft/2020-12/schema" assert "query" in schema["properties"] @@ -137,33 +134,33 @@ def test_chat_mode_includes_query() -> None: app.app_model_config_with_session.return_value.to_dict.assert_called_once_with(annotation_reply={"enabled": False}) -def test_agent_chat_mode_includes_query() -> None: +def test_agent_chat_mode_includes_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.AGENT_CHAT, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" in schema["properties"] -def test_advanced_chat_mode_includes_query() -> None: +def test_advanced_chat_mode_includes_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.ADVANCED_CHAT, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" in schema["properties"] -def test_workflow_mode_omits_query() -> None: +def test_workflow_mode_omits_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.WORKFLOW, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" not in schema["properties"] assert schema["required"] == ["inputs"] -def test_completion_mode_omits_query() -> None: +def test_completion_mode_omits_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.COMPLETION, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" not in schema["properties"] assert schema["required"] == ["inputs"] -def test_inputs_required_driven_by_form() -> None: +def test_inputs_required_driven_by_form(sqlite_session: Session) -> None: app = _stub_app( AppMode.CHAT, form=[ @@ -171,20 +168,20 @@ def test_inputs_required_driven_by_form() -> None: {"text-input": {"variable": "context", "label": "Context", "required": False}}, ], ) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert schema["properties"]["inputs"]["required"] == ["industry"] -def test_misconfigured_chat_raises_app_unavailable() -> None: +def test_misconfigured_chat_raises_app_unavailable(sqlite_session: Session) -> None: app = _stub_app(AppMode.CHAT, has_workflow=False) with pytest.raises(AppUnavailableError): - build_input_schema(app, session=_session()) + build_input_schema(app, session=sqlite_session) -def test_misconfigured_workflow_raises_app_unavailable() -> None: +def test_misconfigured_workflow_raises_app_unavailable(sqlite_session: Session) -> None: app = _stub_app(AppMode.WORKFLOW, has_workflow=False) with pytest.raises(AppUnavailableError): - build_input_schema(app, session=_session()) + build_input_schema(app, session=sqlite_session) def test_empty_input_schema_sentinel_shape() -> None: diff --git a/api/tests/unit_tests/controllers/service_api/app/test_audio.py b/api/tests/unit_tests/controllers/service_api/app/test_audio.py index d41bfc69109..4fb56e6ac8b 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_audio.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_audio.py @@ -160,7 +160,7 @@ class TestAudioServiceMockedBehavior: return mock @patch.object(AudioService, "transcript_asr") - def test_transcript_asr_returns_response(self, mock_asr, mock_app, mock_file): + def test_transcript_asr_returns_response(self, mock_asr, mock_app, mock_file, sqlite_session: Session): """Test ASR transcription returns response dict.""" mock_response = {"text": "Transcribed text"} mock_asr.return_value = mock_response @@ -168,14 +168,13 @@ class TestAudioServiceMockedBehavior: result = AudioService.transcript_asr( app_model=mock_app, file=mock_file, - session=Mock(), + session=sqlite_session, end_user="user_123", ) assert result["text"] == "Transcribed text" @patch.object(AudioService, "transcript_tts") - @pytest.mark.parametrize("sqlite_session", [()], indirect=True) def test_transcript_tts_returns_response(self, mock_tts, mock_app, sqlite_session: Session): """Test TTS transcription returns response.""" mock_response = {"audio": "base64_audio_data"} diff --git a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py index 537e9cda7fe..ac79b565f7a 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py @@ -494,7 +494,6 @@ class TestConversationService: assert hasattr(result, "limit") assert hasattr(result, "has_more") - @pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True) def test_rename_returns_conversation(self, sqlite_session: Session): """Test rename returns updated conversation.""" conversation_id = "00000000-0000-0000-0000-000000000001" @@ -542,7 +541,6 @@ class TestConversationApiController: with pytest.raises(NotChatAppError): handler(api, app_model=app_model, end_user=end_user) - @pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True) def test_list_last_not_found( self, app: Flask, diff --git a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py index dd158731055..57c5c038424 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py @@ -615,7 +615,6 @@ class TestHitlServiceApi: assert response.data.paused_nodes == ["node-1"] assert response.data.reasons == [{"TYPE": "human_input_required", "form_id": "form-1", "expiration_time": 1}] - @pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) def test_service_api_pause_event_serializes_hitl_reason( self, monkeypatch: pytest.MonkeyPatch, @@ -630,7 +629,7 @@ class TestHitlServiceApi: reason=WorkflowStartReason.INITIAL, ) - expiration_time = datetime(2024, 1, 1) + expiration_time = datetime(2024, 1, 1, tzinfo=UTC) _persist_human_input_form(sqlite_session, expiration_time=expiration_time) monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=sqlite_engine)) @@ -693,7 +692,6 @@ class TestHitlServiceApi: assert hi_resp.data.expiration_time == int(expiration_time.timestamp()) # Snapshot payload contract - @pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) def test_snapshot_events_include_pause_payload_contract( self, monkeypatch: pytest.MonkeyPatch, @@ -703,7 +701,7 @@ class TestHitlServiceApi: workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED) snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED) resumption_context = _build_resumption_context("task-ctx") - expiration_time = datetime(2024, 1, 1) + expiration_time = datetime(2024, 1, 1, tzinfo=UTC) _persist_human_input_form(sqlite_session, expiration_time=expiration_time) monkeypatch.setattr( "services.workflow_event_snapshot_service.load_form_dispositions_by_form_id", diff --git a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py index 92018a8c233..d7c524e6532 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py @@ -574,7 +574,6 @@ class TestPipelineRunApiPost: ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - @pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) def test_post_success_streaming( self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app, sqlite_session: Session ): @@ -609,7 +608,6 @@ class TestPipelineRunApiPost: mock_svc_cls.assert_called_once_with(sqlite_session) mock_gen_svc.generate.assert_called_once() - @pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) def test_post_not_found(self, app: Flask, sqlite_session: Session): """Test NotFound when dataset check fails.""" with app.test_request_context("/datasets/test/pipeline/run", method="POST"): @@ -624,7 +622,6 @@ class TestPipelineRunApiPost: @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", new="not_account") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - @pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) def test_post_forbidden_non_account_user(self, mock_ns, app: Flask, sqlite_session: Session): """Test Forbidden when current_user is not an Account.""" tenant_id = str(uuid.uuid4()) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py index c73ef67af85..0452591b88b 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py @@ -20,10 +20,11 @@ import json import uuid from dataclasses import dataclass from datetime import UTC, datetime -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound from controllers.common.errors import FileTooLargeError as FileTooLargeHTTPError @@ -44,8 +45,8 @@ from controllers.service_api.dataset.document import ( ) from controllers.service_api.dataset.error import ArchivedDocumentImmutableError from core.rag.index_processor.constant.index_type import IndexStructureType -from models.dataset import Dataset, Document -from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus, SegmentStatus from services.dataset_service import DocumentService from services.entities.knowledge_entities.knowledge_entities import ProcessRule, RetrievalModel from services.errors.file import FileTooLargeError as FileTooLargeServiceError @@ -176,6 +177,26 @@ def _expected_document_response(document: Document) -> dict[str, object]: } +def _persist_segments(session: Session, document: Document, count: int = 5) -> None: + session.add_all( + [ + DocumentSegment( + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, + position=index, + content=f"segment {index}", + word_count=20, + tokens=5, + created_by=document.created_by, + status=SegmentStatus.COMPLETED, + ) + for index in range(1, count + 1) + ] + ) + session.flush() + + class TestDocumentTextCreatePayload: """Test suite for DocumentTextCreatePayload Pydantic model.""" @@ -370,10 +391,10 @@ class TestDocumentService: assert result.indexing_status == "completed" @patch.object(DocumentService, "delete_document") - def test_delete_document_called(self, mock_delete): + def test_delete_document_called(self, mock_delete, sqlite_session: Session): """Test delete_document is called with document.""" document = make_serializable_document() - session = Mock() + session = sqlite_session DocumentService.delete_document(document=document, session=session) mock_delete.assert_called_once_with(document=document, session=session) @@ -550,24 +571,23 @@ class TestDocumentDisplayStatusLogic: class TestDocumentServiceBatchMethods: """Test DocumentService batch operations.""" - def test_get_documents_by_ids(self): + def test_get_documents_by_ids(self, sqlite_session: Session): """Test batch retrieval of documents by IDs.""" dataset_id = str(uuid.uuid4()) doc_ids = [str(uuid.uuid4()), str(uuid.uuid4())] - session = Mock() - mock_result = Mock() - mock_result.all.return_value = [Mock(id=doc_ids[0]), Mock(id=doc_ids[1])] - session.scalars.return_value = mock_result + session = sqlite_session + session.add_all([make_serializable_document(id=document_id, dataset_id=dataset_id) for document_id in doc_ids]) + session.flush() documents = DocumentService.get_documents_by_ids(dataset_id, doc_ids, session) assert len(documents) == 2 - session.scalars.assert_called_once() + assert {document.id for document in documents} == set(doc_ids) - def test_get_documents_by_ids_empty(self): + def test_get_documents_by_ids_empty(self, sqlite_session: Session): """Test batch retrieval with empty list returns empty.""" - assert DocumentService.get_documents_by_ids("ds_id", [], Mock()) == [] + assert DocumentService.get_documents_by_ids("ds_id", [], sqlite_session) == [] class TestDocumentServiceFileOperations: @@ -575,7 +595,7 @@ class TestDocumentServiceFileOperations: @patch("services.dataset_service.file_helpers.get_signed_file_url") @patch("services.dataset_service.DocumentService._get_upload_file_for_upload_file_document") - def test_get_document_download_url(self, mock_get_file, mock_signed_url): + def test_get_document_download_url(self, mock_get_file, mock_signed_url, sqlite_session: Session): """Test generation of download URL.""" mock_doc = Mock() mock_file = Mock() @@ -583,7 +603,7 @@ class TestDocumentServiceFileOperations: mock_get_file.return_value = mock_file mock_signed_url.return_value = "https://example.com/download" - session = Mock() + session = sqlite_session url = DocumentService.get_document_download_url(mock_doc, session) assert url == "https://example.com/download" @@ -597,7 +617,7 @@ class TestDocumentServiceSaveValidation: @patch("services.dataset_service.DatasetService.check_doc_form") @patch("services.dataset_service.FeatureService.get_features") @patch("services.dataset_service.current_user") - def test_save_document_validates_doc_form(self, mock_user, mock_features, mock_check_form): + def test_save_document_validates_doc_form(self, mock_user, mock_features, mock_check_form, sqlite_session: Session): """Test that doc_form is validated during save.""" mock_user.current_tenant_id = "tenant_id" dataset = Mock() @@ -610,7 +630,7 @@ class TestDocumentServiceSaveValidation: pass mock_check_form.side_effect = TestStopError() - session = Mock() + session = sqlite_session # Skip actual logic by mocking dependent calls or raising error to stop early with pytest.raises(TestStopError): @@ -663,6 +683,7 @@ class TestDocumentApiGet: app: Flask, mock_tenant: str, mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test successful document retrieval with metadata='all'.""" # Arrange @@ -670,9 +691,10 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}} - session = MagicMock() - session.scalar.side_effect = [5, 0] + session = sqlite_session + _persist_segments(session, mock_doc_detail) # Act with app.test_request_context( @@ -726,15 +748,16 @@ class TestDocumentApiGet: assert response["summary_index_status"] is None @patch("controllers.service_api.dataset.document.DocumentService") - def test_get_document_not_found(self, mock_doc_svc: Mock, app: Flask, mock_tenant: str) -> None: + def test_get_document_not_found( + self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, sqlite_session: Session + ) -> None: """Test 404 when document is not found.""" # Arrange dataset_id = str(uuid.uuid4()) mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant) mock_doc_svc.get_document.return_value = None - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session # Act & Assert with app.test_request_context( @@ -754,7 +777,12 @@ class TestDocumentApiGet: @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_forbidden_wrong_tenant( - self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document + self, + mock_doc_svc: Mock, + app: Flask, + mock_tenant: str, + mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test 403 when document tenant doesn't match request tenant.""" # Arrange @@ -763,8 +791,9 @@ class TestDocumentApiGet: mock_doc_detail.tenant_id = "different-tenant-id" mock_doc_svc.get_document.return_value = mock_doc_detail - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session + session.add(mock_dataset) + session.flush() # Act & Assert with app.test_request_context( @@ -784,7 +813,12 @@ class TestDocumentApiGet: @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_metadata_only( - self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document + self, + mock_doc_svc: Mock, + app: Flask, + mock_tenant: str, + mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test document retrieval with metadata='only'.""" # Arrange @@ -792,8 +826,9 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session + session.add(mock_dataset) + session.flush() # Act with app.test_request_context( @@ -827,6 +862,7 @@ class TestDocumentApiGet: app: Flask, mock_tenant: str, mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test document retrieval with metadata='without'.""" # Arrange @@ -834,9 +870,10 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}} - session = MagicMock() - session.scalar.side_effect = [5, 0] + session = sqlite_session + _persist_segments(session, mock_doc_detail) # Act with app.test_request_context( @@ -895,7 +932,12 @@ class TestDocumentApiGet: @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_invalid_metadata_value( - self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document + self, + mock_doc_svc: Mock, + app: Flask, + mock_tenant: str, + mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test error when metadata parameter has invalid value.""" # Arrange @@ -903,8 +945,9 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session + session.add(mock_dataset) + session.flush() # Act & Assert with app.test_request_context( @@ -1464,6 +1507,7 @@ class TestDocumentUpdateByTextApiPost: app: Flask, mock_tenant, mock_dataset, + sqlite_session: Session, ): """Test successful document update by text.""" _setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant) @@ -1799,6 +1843,7 @@ class TestDocumentUpdateByFileApiPatch: app: Flask, mock_tenant, mock_dataset, + sqlite_session: Session, ): """Test legacy POST aliases still dispatch while marked deprecated.""" _setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant) @@ -1808,8 +1853,7 @@ class TestDocumentUpdateByFileApiPatch: ) doc_id = str(uuid.uuid4()) - session = MagicMock() - session.scalar.return_value = 0 + session = sqlite_session with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/{doc_id}/{route_name}", method="POST", diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py index c30762911e2..af6d0b3699c 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py @@ -17,10 +17,11 @@ Decorator strategy: import uuid from inspect import unwrap -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.service_api.dataset.metadata import ( @@ -56,7 +57,15 @@ def mock_dataset(): # --------------------------------------------------------------------------- -class TestDatasetMetadataCreatePost: +class _UsesSQLiteSession: + session: Session + + @pytest.fixture(autouse=True) + def _inject_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + + +class TestDatasetMetadataCreatePost(_UsesSQLiteSession): """Tests for DatasetMetadataCreateServiceApi.post(). ``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check`` @@ -64,7 +73,7 @@ class TestDatasetMetadataCreatePost: """ @staticmethod - def _call_post(api, session: MagicMock, **kwargs): + def _call_post(api, session: Session, **kwargs): return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -91,7 +100,7 @@ class TestDatasetMetadataCreatePost: json={"type": "string", "name": "Author"}, ): api = DatasetMetadataCreateServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -120,7 +129,7 @@ class TestDatasetMetadataCreatePost: json={"type": "string", "name": "Author"}, ): api = DatasetMetadataCreateServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_post( api, @@ -130,7 +139,7 @@ class TestDatasetMetadataCreatePost: ) -class TestDatasetMetadataCreateGet: +class TestDatasetMetadataCreateGet(_UsesSQLiteSession): """Tests for DatasetMetadataCreateServiceApi.get().""" @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -191,14 +200,14 @@ class TestDatasetMetadataCreateGet: # --------------------------------------------------------------------------- -class TestDatasetMetadataServiceApiPatch: +class TestDatasetMetadataServiceApiPatch(_UsesSQLiteSession): """Tests for DatasetMetadataServiceApi.patch(). ``patch`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_patch(api, session: MagicMock, **kwargs): + def _call_patch(api, session: Session, **kwargs): return unwrap(api.patch)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -225,7 +234,7 @@ class TestDatasetMetadataServiceApiPatch: json={"name": "New Name"}, ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session response, status = self._call_patch( api, session, @@ -256,7 +265,7 @@ class TestDatasetMetadataServiceApiPatch: json={"name": "x"}, ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_patch( api, @@ -267,14 +276,14 @@ class TestDatasetMetadataServiceApiPatch: ) -class TestDatasetMetadataServiceApiDelete: +class TestDatasetMetadataServiceApiDelete(_UsesSQLiteSession): """Tests for DatasetMetadataServiceApi.delete(). ``delete`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_delete(api, session: MagicMock, **kwargs): + def _call_delete(api, session: Session, **kwargs): return unwrap(api.delete)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -300,7 +309,7 @@ class TestDatasetMetadataServiceApiDelete: method="DELETE", ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session response = self._call_delete( api, session, @@ -329,7 +338,7 @@ class TestDatasetMetadataServiceApiDelete: method="DELETE", ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_delete( api, @@ -345,7 +354,7 @@ class TestDatasetMetadataServiceApiDelete: # --------------------------------------------------------------------------- -class TestDatasetMetadataBuiltInFieldGet: +class TestDatasetMetadataBuiltInFieldGet(_UsesSQLiteSession): """Tests for DatasetMetadataBuiltInFieldServiceApi.get().""" @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -380,14 +389,14 @@ class TestDatasetMetadataBuiltInFieldGet: # --------------------------------------------------------------------------- -class TestDatasetMetadataBuiltInFieldAction: +class TestDatasetMetadataBuiltInFieldAction(_UsesSQLiteSession): """Tests for DatasetMetadataBuiltInFieldActionServiceApi.post(). ``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_post(api, session: MagicMock, **kwargs): + def _call_post(api, session: Session, **kwargs): return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -411,7 +420,7 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -445,7 +454,7 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -473,7 +482,7 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_post( api, @@ -489,14 +498,14 @@ class TestDatasetMetadataBuiltInFieldAction: # --------------------------------------------------------------------------- -class TestDocumentMetadataEditPost: +class TestDocumentMetadataEditPost(_UsesSQLiteSession): """Tests for DocumentMetadataEditServiceApi.post(). ``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_post(api, session: MagicMock, **kwargs): + def _call_post(api, session: Session, **kwargs): return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -522,7 +531,7 @@ class TestDocumentMetadataEditPost: json={"operation_data": []}, ): api = DocumentMetadataEditServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -550,7 +559,7 @@ class TestDocumentMetadataEditPost: json={"operation_data": []}, ): api = DocumentMetadataEditServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_post( api, diff --git a/api/tests/unit_tests/controllers/web/test_completion.py b/api/tests/unit_tests/controllers/web/test_completion.py index e3bbbe2c87c..e76a4d7d6a4 100644 --- a/api/tests/unit_tests/controllers/web/test_completion.py +++ b/api/tests/unit_tests/controllers/web/test_completion.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from controllers.web.completion import ChatApi, ChatStopApi, CompletionApi, CompletionStopApi from controllers.web.error import ( @@ -168,9 +169,10 @@ class TestChatApi: mock_get_conversation: MagicMock, mock_generate: MagicMock, app: Flask, + unbound_session: Session, ) -> None: mock_ns.payload = {"inputs": {}, "query": "hi", "conversation_id": str(uuid.uuid4())} - session = MagicMock() + session = unbound_session with app.test_request_context("/chat-messages", method="POST"): unwrap(ChatApi.post)(ChatApi(), session, _chat_app(), _end_user())