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 14cf4584851..436d9469992 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py @@ -1,4 +1,5 @@ import logging +from collections.abc import Iterator from http import HTTPStatus from inspect import unwrap from io import BytesIO @@ -6,6 +7,8 @@ from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, scoped_session, sessionmaker from werkzeug.datastructures import FileStorage from werkzeug.exceptions import Unauthorized @@ -36,6 +39,17 @@ from libs.datetime_utils import naive_utc_now from models.account import Account, Tenant, TenantCustomConfigDict, TenantStatus +@pytest.fixture +def workspace_session(sqlite_engine: Engine) -> Iterator[scoped_session[Session]]: + """Provide the callable scoped session expected by Flask-SQLAlchemy controllers.""" + Tenant.metadata.create_all(sqlite_engine, tables=[Tenant.__table__]) + session = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False)) + try: + yield session + finally: + session.remove() + + def make_account(account_id: str = "u1") -> Account: account = Account(name="Test User", email=f"{account_id}@example.com") account.id = account_id @@ -370,11 +384,13 @@ class TestTenantInfoResponse: class TestSwitchWorkspaceApi: - def test_switch_success(self, app: Flask): + def test_switch_success(self, app: Flask, workspace_session: scoped_session[Session]): api = SwitchWorkspaceApi() method = unwrap(api.post) payload = {"tenant_id": "t2"} tenant = make_tenant("t2") + workspace_session.add(tenant) + workspace_session.commit() user = make_account() with ( app.test_request_context("/workspaces/switch", json=payload), @@ -383,11 +399,10 @@ class TestSwitchWorkspaceApi: "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t2"} ), ): - session = MagicMock() - session.get.return_value = tenant - result = method(api, session, user) + result = method(api, workspace_session, user) + assert result["result"] == "success" - switch_tenant.assert_called_once_with(user, "t2", session=session) + switch_tenant.assert_called_once_with(user, "t2", session=workspace_session) def test_switch_not_linked(self, app: Flask): api = SwitchWorkspaceApi() @@ -401,7 +416,7 @@ class TestSwitchWorkspaceApi: with pytest.raises(AccountNotLinkTenantError): method(api, MagicMock(), user) - def test_switch_tenant_not_found(self, app: Flask): + def test_switch_tenant_not_found(self, app: Flask, workspace_session: scoped_session[Session]): api = SwitchWorkspaceApi() method = unwrap(api.post) payload = {"tenant_id": "missing"} @@ -410,19 +425,21 @@ class TestSwitchWorkspaceApi: app.test_request_context("/workspaces/switch", json=payload), patch("controllers.console.workspace.workspace.TenantService.switch_tenant"), ): - session = MagicMock() - session.get.return_value = None with pytest.raises(ValueError): - method(api, session, user) + method(api, workspace_session, user) class TestCustomConfigWorkspaceApi: - def test_post_success(self, app: Flask): + def test_post_success(self, app: Flask, workspace_session: scoped_session[Session]): api = CustomConfigWorkspaceApi() method = unwrap(api.post) tenant = make_tenant(custom_config={}) + workspace_session.add(tenant) + workspace_session.commit() + payload = {"remove_webapp_brand": True} events = [] + event.listen(workspace_session, "after_commit", lambda _: events.append("commit")) with ( app.test_request_context("/workspaces/custom-config", json=payload), patch( @@ -430,27 +447,29 @@ class TestCustomConfigWorkspaceApi: side_effect=lambda *args, **kwargs: events.append("get_tenant_info") or {"id": "t1"}, ), ): - session = MagicMock() - session.get.return_value = tenant - session.commit.side_effect = lambda: events.append("commit") - result = method(api, session, "t1") + result = method(api, workspace_session, "t1") assert result["result"] == "success" assert events == ["commit", "get_tenant_info"] - def test_logo_fallback(self, app: Flask): + def test_logo_fallback(self, app: Flask, workspace_session: scoped_session[Session]): api = CustomConfigWorkspaceApi() method = unwrap(api.post) + tenant = make_tenant(custom_config={"replace_webapp_logo": "old-logo"}) + workspace_session.add(tenant) + workspace_session.commit() + payload = {"remove_webapp_brand": False} + with ( app.test_request_context("/workspaces/custom-config", json=payload), patch( - "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t1"} + "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", + return_value={"id": "t1"}, ), ): - session = MagicMock() - session.get.return_value = tenant - result = method(api, session, "t1") + result = method(api, workspace_session, "t1") + assert tenant.custom_config_dict["replace_webapp_logo"] == "old-logo" assert result["result"] == "success" @@ -541,14 +560,19 @@ class TestWebappLogoWorkspaceApi: class TestWorkspaceInfoApi: - def test_post_success(self, app: Flask): + def test_post_success(self, app: Flask, workspace_session: scoped_session[Session]): api = WorkspaceInfoApi() method = unwrap(api.post) tenant = make_tenant() + workspace_session.add(tenant) + workspace_session.commit() + payload = {"name": "New Name"} events = [] with ( app.test_request_context("/workspaces/info", json=payload), + patch("controllers.console.workspace.workspace.db.get_or_404", return_value=tenant), + patch("controllers.console.workspace.workspace.db.session", workspace_session), patch( "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", side_effect=lambda *args, **kwargs: ( diff --git a/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py b/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py index 9dece8cfe07..71b17c96698 100644 --- a/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py +++ b/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py @@ -3,7 +3,10 @@ from __future__ import annotations from types import SimpleNamespace -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch + +import pytest +from sqlalchemy.orm import Session from core.app.app_config.entities import DatasetRetrieveConfigEntity from core.app.entities.app_invoke_entities import InvokeFrom @@ -14,13 +17,14 @@ def _retrieve_config() -> DatasetRetrieveConfigEntity: return DatasetRetrieveConfigEntity(retrieve_strategy=DatasetRetrieveConfigEntity.RetrieveStrategy.MULTIPLE) -def test_get_dataset_tools_returns_empty_for_empty_dataset_ids() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_get_dataset_tools_returns_empty_for_empty_dataset_ids(sqlite_session: Session) -> None: # Arrange retrieve_config = _retrieve_config() # Act tools = DatasetRetrieverTool.get_dataset_tools( - session=MagicMock(), + session=sqlite_session, tenant_id="tenant", dataset_ids=[], retrieve_config=retrieve_config, @@ -35,13 +39,14 @@ def test_get_dataset_tools_returns_empty_for_empty_dataset_ids() -> None: assert tools == [] -def test_get_dataset_tools_returns_empty_for_missing_retrieve_config() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_get_dataset_tools_returns_empty_for_missing_retrieve_config(sqlite_session: Session) -> None: # Arrange dataset_ids = ["d1"] # Act tools = DatasetRetrieverTool.get_dataset_tools( - session=MagicMock(), + session=sqlite_session, tenant_id="tenant", dataset_ids=dataset_ids, retrieve_config=None, # type: ignore[arg-type] @@ -56,7 +61,8 @@ def test_get_dataset_tools_returns_empty_for_missing_retrieve_config() -> None: assert tools == [] -def test_get_dataset_tools_builds_tool_and_restores_strategy() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_get_dataset_tools_builds_tool_and_restores_strategy(sqlite_session: Session) -> None: # Arrange retrieve_config = _retrieve_config() retrieval_tool = SimpleNamespace(name="dataset_tool", description="desc", run=lambda query: f"result:{query}") @@ -66,7 +72,7 @@ def test_get_dataset_tools_builds_tool_and_restores_strategy() -> None: # Act with patch("core.tools.utils.dataset_retriever_tool.DatasetRetrieval", return_value=feature): tools = DatasetRetrieverTool.get_dataset_tools( - session=MagicMock(), + session=sqlite_session, tenant_id="tenant", dataset_ids=["d1"], retrieve_config=retrieve_config, @@ -83,7 +89,7 @@ def test_get_dataset_tools_builds_tool_and_restores_strategy() -> None: assert retrieve_config.retrieve_strategy == DatasetRetrieveConfigEntity.RetrieveStrategy.MULTIPLE -def _build_dataset_tool() -> tuple[DatasetRetrieverTool, SimpleNamespace]: +def _build_dataset_tool(sqlite_session: Session) -> tuple[DatasetRetrieverTool, SimpleNamespace]: retrieval_tool = SimpleNamespace( name="dataset_tool", description="desc", @@ -93,7 +99,7 @@ def _build_dataset_tool() -> tuple[DatasetRetrieverTool, SimpleNamespace]: feature.to_dataset_retriever_tool.return_value = [retrieval_tool] with patch("core.tools.utils.dataset_retriever_tool.DatasetRetrieval", return_value=feature): tools = DatasetRetrieverTool.get_dataset_tools( - session=MagicMock(), + session=sqlite_session, tenant_id="tenant", dataset_ids=["d1"], retrieve_config=_retrieve_config(), @@ -106,9 +112,10 @@ def _build_dataset_tool() -> tuple[DatasetRetrieverTool, SimpleNamespace]: return tools[0], retrieval_tool -def test_runtime_parameters_shape() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_runtime_parameters_shape(sqlite_session: Session) -> None: # Arrange - tool, _ = _build_dataset_tool() + tool, _ = _build_dataset_tool(sqlite_session) # Act params = tool.get_runtime_parameters() @@ -118,33 +125,36 @@ def test_runtime_parameters_shape() -> None: assert params[0].name == "query" -def test_empty_query_behavior() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_empty_query_behavior(sqlite_session: Session) -> None: # Arrange - tool, _ = _build_dataset_tool() + tool, _ = _build_dataset_tool(sqlite_session) # Act - empty_query = list(tool.invoke(session=MagicMock(), user_id="u", tool_parameters={})) + empty_query = list(tool.invoke(session=sqlite_session, user_id="u", tool_parameters={})) # Assert assert len(empty_query) == 1 assert empty_query[0].message.text == "please input query" -def test_query_invocation_result() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_query_invocation_result(sqlite_session: Session) -> None: # Arrange - tool, _ = _build_dataset_tool() + tool, _ = _build_dataset_tool(sqlite_session) # Act - result = list(tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"query": "hello"})) + result = list(tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"query": "hello"})) # Assert assert len(result) == 1 assert result[0].message.text == "result:hello" -def test_validate_credentials() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_validate_credentials(sqlite_session: Session) -> None: # Arrange - tool, _ = _build_dataset_tool() + tool, _ = _build_dataset_tool(sqlite_session) # Act result = tool.validate_credentials(credentials={}, parameters={}, format_only=False) diff --git a/api/tests/unit_tests/core/tools/test_mcp_tool.py b/api/tests/unit_tests/core/tools/test_mcp_tool.py index 984794ae967..be0ce20ad60 100644 --- a/api/tests/unit_tests/core/tools/test_mcp_tool.py +++ b/api/tests/unit_tests/core/tools/test_mcp_tool.py @@ -1,9 +1,12 @@ +"""MCP tool tests using real SQLite sessions for the ORM invocation contract.""" + from __future__ import annotations import base64 -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.mcp.types import ( @@ -97,7 +100,8 @@ def test_mcp_tool_usage_extraction_helpers(): assert derived.total_tokens == 0 -def test_mcp_tool_invoke_handles_content_types_and_structured_output(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_mcp_tool_invoke_handles_content_types_and_structured_output(sqlite_session: Session): tool = _build_mcp_tool() img_data = base64.b64encode(b"img").decode() blob_data = base64.b64encode(b"blob").decode() @@ -123,7 +127,7 @@ def test_mcp_tool_invoke_handles_content_types_and_structured_output(): ) with patch.object(MCPTool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"a": 1})) + messages = list(tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={"a": 1})) types = [m.type for m in messages] assert ToolInvokeMessage.MessageType.JSON in types @@ -133,7 +137,8 @@ def test_mcp_tool_invoke_handles_content_types_and_structured_output(): assert tool.latest_usage.total_tokens == 5 -def test_mcp_tool_invoke_raises_for_unsupported_embedded_resource(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_mcp_tool_invoke_raises_for_unsupported_embedded_resource(sqlite_session: Session): tool = _build_mcp_tool() # Use model_construct to bypass pydantic validation and force unsupported resource path. bad_resource = EmbeddedResource.model_construct(type="resource", resource=object()) @@ -141,7 +146,7 @@ def test_mcp_tool_invoke_raises_for_unsupported_embedded_resource(): with patch.object(MCPTool, "invoke_remote_mcp_tool", return_value=result): with pytest.raises(ToolInvokeError, match="Unsupported embedded resource type"): - list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={})) + list(tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={})) def test_mcp_tool_handle_none_parameter_filters_empty_values():