test: use SQLite sessions in unit misc (#39121)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Asuka Minato 2026-07-20 19:38:52 +09:00 committed by GitHub
parent 23882a704e
commit cbd7872520
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 83 additions and 44 deletions

View File

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

View File

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

View File

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