test: use SQLite sessions in core tools (#39068)

This commit is contained in:
Asuka Minato 2026-07-31 21:53:50 +09:00 committed by GitHub
parent 9ee7750411
commit bab3db7f77
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -1,14 +1,20 @@
from __future__ import annotations
"""Unit tests for ToolManager with persisted providers and isolated external collaborators."""
"""Unit tests for ToolManager behavior with mocked providers and collaborators."""
from __future__ import annotations
import json
import threading
from collections.abc import Iterator
from dataclasses import dataclass
from datetime import datetime
from types import SimpleNamespace
from typing import Any
from unittest.mock import MagicMock, Mock, patch
from unittest.mock import Mock, patch
import pytest
from sqlalchemy import event
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom
from core.plugin.entities.plugin_daemon import CredentialType
@ -22,6 +28,94 @@ from core.tools.entities.tool_entities import (
from core.tools.errors import ToolProviderNotFoundError
from core.tools.plugin_tool.provider import PluginToolProviderController
from core.tools.tool_manager import ToolManager
from models.base import TypeBase
from models.tools import ApiToolProvider, BuiltinToolProvider, WorkflowToolProvider
@dataclass(frozen=True)
class _ToolDatabase:
engine: Engine
session: Session
@pytest.fixture
def tool_database(sqlite_engine: Engine) -> Iterator[_ToolDatabase]:
"""Provide isolated provider tables through the same engine/session split used in production."""
tables = [
TypeBase.metadata.tables[model.__tablename__]
for model in (BuiltinToolProvider, ApiToolProvider, WorkflowToolProvider)
]
TypeBase.metadata.create_all(sqlite_engine, tables=tables)
with Session(sqlite_engine, expire_on_commit=False) as session:
yield _ToolDatabase(engine=sqlite_engine, session=session)
def _builtin_provider(
*,
provider_id: str,
tenant_id: str,
provider: str = "time",
name: str = "Time credentials",
encrypted_credentials: str = '{"encrypted":"value"}',
is_default: bool = True,
credential_type: CredentialType = CredentialType.API_KEY,
expires_at: int = -1,
) -> BuiltinToolProvider:
record = BuiltinToolProvider(
tenant_id=tenant_id,
user_id="00000000-0000-0000-0000-000000000099",
provider=provider,
name=name,
encrypted_credentials=encrypted_credentials,
is_default=is_default,
credential_type=credential_type,
expires_at=expires_at,
)
record.id = provider_id
return record
def _api_provider(
*, provider_id: str, tenant_id: str, name: str = "api-provider", icon: str = '{"background":"#000","content":"A"}'
) -> ApiToolProvider:
record = ApiToolProvider(
name=name,
icon=icon,
schema="{}",
schema_type_str="openapi",
user_id="00000000-0000-0000-0000-000000000099",
tenant_id=tenant_id,
description="desc",
tools_str="[]",
credentials_str='{"auth_type":"api_key_query","api_key_value":"secret"}',
privacy_policy="privacy",
custom_disclaimer="disclaimer",
)
record.id = provider_id
return record
def _workflow_provider(
*,
provider_id: str,
tenant_id: str,
name: str = "workflow-provider",
icon: str = '{"background":"#222","content":"W"}',
) -> WorkflowToolProvider:
record = WorkflowToolProvider(
name=name,
label=name,
icon=icon,
app_id=provider_id,
version="1",
user_id="00000000-0000-0000-0000-000000000099",
tenant_id=tenant_id,
description="desc",
parameter_configuration="[]",
)
record.id = provider_id
return record
class _SimpleContextVar:
@ -39,26 +133,24 @@ class _SimpleContextVar:
self._is_set = True
def _cm(session: Any):
context = Mock()
context.__enter__ = Mock(return_value=session)
context.__exit__ = Mock(return_value=False)
return context
def _setup_list_providers_from_api_mocks(
monkeypatch,
*,
session: Mock,
tool_database: _ToolDatabase,
hardcoded_controller: SimpleNamespace,
plugin_controller: PluginToolProviderController,
api_controller: SimpleNamespace,
workflow_controller: SimpleNamespace,
):
mock_db = Mock()
mock_db.engine = object()
monkeypatch.setattr("core.tools.tool_manager.db", mock_db)
monkeypatch.setattr("core.tools.tool_manager.Session", lambda *args, **kwargs: _cm(session))
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
monkeypatch.setattr(
"core.tools.tool_manager.dify_config",
SimpleNamespace(
SQLALCHEMY_DATABASE_URI_SCHEME="mysql",
POSITION_TOOL_INCLUDES_SET=None,
POSITION_TOOL_EXCLUDES_SET=None,
),
)
monkeypatch.setattr(
ToolManager,
"list_builtin_providers",
@ -92,7 +184,12 @@ def _setup_list_providers_from_api_mocks(
)
monkeypatch.setattr(
"core.tools.tool_manager.ToolLabelManager.get_tools_labels",
Mock(side_effect=[{"api-1": ["search"]}, {"wf-1": ["utility"]}]),
Mock(
side_effect=[
{api_controller.provider_id: ["search"]},
{workflow_controller.provider_id: ["utility"]},
]
),
)
mock_mcp_service = Mock()
mock_mcp_service.list_providers.return_value = [SimpleNamespace(name="mcp-provider")]
@ -201,7 +298,9 @@ def test_get_tool_runtime_builtin_missing_tool_raises():
)
def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks():
def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks(
monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase
):
tool = Mock()
tool.fork_tool_runtime.return_value = "runtime-tool"
controller = SimpleNamespace(
@ -209,28 +308,27 @@ def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks():
need_credentials=True,
get_credentials_schema_by_type=Mock(return_value=[]),
)
builtin_provider = SimpleNamespace(
id="cred-1",
credential_type=CredentialType.API_KEY.value,
credentials={"encrypted": "value"},
expires_at=-1,
user_id="user-1",
tenant_id = "00000000-0000-0000-0000-000000000001"
builtin_provider = _builtin_provider(
provider_id="00000000-0000-0000-0000-000000000002",
tenant_id=tenant_id,
)
tool_database.session.add(builtin_provider)
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
with patch.object(ToolManager, "get_builtin_provider", return_value=controller):
with patch("core.helper.credential_utils.check_credential_policy_compliance"):
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.session.scalar.return_value = builtin_provider
encrypter = Mock()
encrypter.decrypt.return_value = {"api_key": "secret"}
cache = Mock()
with patch("core.tools.tool_manager.create_provider_encrypter", return_value=(encrypter, cache)):
result = ToolManager.get_tool_runtime(
provider_type=ToolProviderType.BUILT_IN,
provider_id="time",
tool_name="weekday",
tenant_id="tenant-1",
)
encrypter = Mock()
encrypter.decrypt.return_value = {"api_key": "secret"}
cache = Mock()
with patch("core.tools.tool_manager.create_provider_encrypter", return_value=(encrypter, cache)):
result = ToolManager.get_tool_runtime(
provider_type=ToolProviderType.BUILT_IN,
provider_id="time",
tool_name="weekday",
tenant_id=tenant_id,
)
assert result == "runtime-tool"
runtime = tool.fork_tool_runtime.call_args.kwargs["runtime"]
@ -244,16 +342,16 @@ def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks():
"services.tools.builtin_tools_manage_service.BuiltinToolManageService.get_oauth_client",
return_value={"client_id": "id"},
)
@patch("core.tools.tool_manager.db")
@patch("core.tools.tool_manager.time.time", return_value=1000)
@patch("core.helper.credential_utils.check_credential_policy_compliance")
def test_get_tool_runtime_builtin_refreshes_expired_oauth_credentials(
mock_check,
mock_time,
mock_db,
mock_get_oauth_client,
mock_oauth_handler_cls,
mock_create_provider_encrypter,
monkeypatch: pytest.MonkeyPatch,
tool_database: _ToolDatabase,
):
tool = Mock()
tool.fork_tool_runtime.return_value = "runtime-tool"
@ -262,17 +360,19 @@ def test_get_tool_runtime_builtin_refreshes_expired_oauth_credentials(
need_credentials=True,
get_credentials_schema_by_type=Mock(return_value=[]),
)
builtin_provider = SimpleNamespace(
id="cred-1",
credential_type=CredentialType.OAUTH2.value,
credentials={"encrypted": "value"},
encrypted_credentials=None,
tenant_id = "00000000-0000-0000-0000-000000000001"
provider_id = "00000000-0000-0000-0000-000000000002"
builtin_provider = _builtin_provider(
provider_id=provider_id,
tenant_id=tenant_id,
credential_type=CredentialType.OAUTH2,
expires_at=1,
user_id="user-1",
)
refreshed = SimpleNamespace(credentials={"token": "new"}, expires_at=123456)
mock_db.session.scalar.return_value = builtin_provider
tool_database.session.add(builtin_provider)
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
encrypter = Mock()
encrypter.decrypt.return_value = {"token": "old"}
encrypter.encrypt.return_value = {"token": "encrypted"}
@ -285,34 +385,38 @@ def test_get_tool_runtime_builtin_refreshes_expired_oauth_credentials(
provider_type=ToolProviderType.BUILT_IN,
provider_id="time",
tool_name="weekday",
tenant_id="tenant-1",
tenant_id=tenant_id,
)
assert result == "runtime-tool"
assert builtin_provider.expires_at == refreshed.expires_at
assert builtin_provider.encrypted_credentials == json.dumps({"token": "encrypted"})
mock_db.session.commit.assert_called_once()
tool_database.session.expire_all()
persisted = tool_database.session.get(BuiltinToolProvider, provider_id)
assert persisted is not None
assert persisted.expires_at == refreshed.expires_at
assert persisted.encrypted_credentials == json.dumps({"token": "encrypted"})
cache.delete.assert_called_once()
def test_get_tool_runtime_builtin_plugin_provider_deleted_raises():
def test_get_tool_runtime_builtin_plugin_provider_deleted_raises(
monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase
):
plugin_controller = object.__new__(PluginToolProviderController)
plugin_controller.entity = SimpleNamespace(credentials_schema=[{"name": "k"}], oauth_schema=None)
plugin_controller.get_tool = Mock(return_value=Mock())
plugin_controller.get_credentials_schema_by_type = Mock(return_value=[])
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
with patch.object(ToolManager, "get_builtin_provider", return_value=plugin_controller):
with patch("core.tools.tool_manager.is_valid_uuid", return_value=True):
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.session.scalar.return_value = None
with pytest.raises(ToolProviderNotFoundError, match="provider has been deleted"):
ToolManager.get_tool_runtime(
provider_type=ToolProviderType.BUILT_IN,
provider_id="time",
tool_name="weekday",
tenant_id="tenant-1",
credential_id="uuid-id",
)
with pytest.raises(ToolProviderNotFoundError, match="provider has been deleted"):
ToolManager.get_tool_runtime(
provider_type=ToolProviderType.BUILT_IN,
provider_id="time",
tool_name="weekday",
tenant_id="00000000-0000-0000-0000-000000000001",
credential_id="00000000-0000-0000-0000-000000000002",
)
def test_get_tool_runtime_api_path():
@ -336,32 +440,30 @@ def test_get_tool_runtime_api_path():
)
def test_get_tool_runtime_workflow_path():
workflow_provider = SimpleNamespace(tenant_id="tenant-1")
def test_get_tool_runtime_workflow_path(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase):
tenant_id = "00000000-0000-0000-0000-000000000001"
provider_id = "00000000-0000-0000-0000-000000000002"
workflow_provider = _workflow_provider(provider_id=provider_id, tenant_id=tenant_id)
tool_database.session.add(workflow_provider)
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
workflow_tool = Mock()
workflow_tool.fork_tool_runtime.return_value = "wf-runtime"
workflow_controller = Mock()
workflow_controller.get_tools.return_value = [workflow_tool]
session = Mock()
session.begin.return_value = _cm(None)
session.scalar.return_value = workflow_provider
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.engine = object()
with patch("core.tools.tool_manager.Session", return_value=_cm(session)):
with patch(
"core.tools.tool_manager.ToolTransformService.workflow_provider_to_controller",
return_value=workflow_controller,
):
assert (
ToolManager.get_tool_runtime(
provider_type=ToolProviderType.WORKFLOW,
provider_id="wf-1",
tool_name="wf",
tenant_id="tenant-1",
)
== "wf-runtime"
)
with patch(
"core.tools.tool_manager.ToolTransformService.workflow_provider_to_controller",
return_value=workflow_controller,
):
assert (
ToolManager.get_tool_runtime(
provider_type=ToolProviderType.WORKFLOW,
provider_id=provider_id,
tool_name="wf",
tenant_id=tenant_id,
)
== "wf-runtime"
)
def test_get_tool_runtime_plugin_path():
@ -631,79 +733,152 @@ def test_get_tool_label_loads_cache_and_handles_missing():
assert ToolManager.get_tool_label("missing") is None
def test_list_default_builtin_providers_for_postgres_and_mysql():
provider_records = [SimpleNamespace(id="id-1"), SimpleNamespace(id="id-2")]
@pytest.mark.parametrize("database_scheme", ["mysql", "postgresql"])
def test_list_default_builtin_providers_uses_persisted_defaults(
monkeypatch: pytest.MonkeyPatch,
tool_database: _ToolDatabase,
database_scheme: str,
):
tenant_id = "00000000-0000-0000-0000-000000000001"
default_provider = _builtin_provider(
provider_id="00000000-0000-0000-0000-000000000002",
tenant_id=tenant_id,
name="default",
is_default=True,
)
older_provider = _builtin_provider(
provider_id="00000000-0000-0000-0000-000000000003",
tenant_id=tenant_id,
name="older",
is_default=False,
)
other_tenant_provider = _builtin_provider(
provider_id="00000000-0000-0000-0000-000000000004",
tenant_id="00000000-0000-0000-0000-000000000005",
name="foreign",
)
default_provider.created_at = datetime(2026, 1, 2)
older_provider.created_at = datetime(2026, 1, 1)
tool_database.session.add_all([default_provider, older_provider, other_tenant_provider])
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
monkeypatch.setattr(
"core.tools.tool_manager.dify_config",
SimpleNamespace(SQLALCHEMY_DATABASE_URI_SCHEME=database_scheme),
)
for scheme in ("postgresql", "mysql"):
session = Mock()
session.execute.return_value.all.return_value = [SimpleNamespace(id="id-1"), SimpleNamespace(id="id-2")]
session.scalars.return_value = iter(provider_records)
postgresql_statements = []
with patch("core.tools.tool_manager.dify_config", SimpleNamespace(SQLALCHEMY_DATABASE_URI_SCHEME=scheme)):
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.engine = object()
with patch("core.tools.tool_manager.Session", return_value=_cm(session)):
providers = ToolManager.list_default_builtin_providers("tenant-1")
def translate_postgresql_distinct_on(_connection, _cursor, statement, parameters, _context, _executemany):
if "SELECT DISTINCT ON (tenant_id, provider) id" not in statement:
return statement, parameters
assert providers == provider_records
postgresql_statements.append(statement)
sqlite_statement = """
SELECT id FROM (
SELECT id,
ROW_NUMBER() OVER (
PARTITION BY tenant_id, provider
ORDER BY is_default DESC, created_at DESC
) AS rn
FROM tool_builtin_providers
WHERE tenant_id = ?
) ranked WHERE rn = 1
"""
return sqlite_statement, parameters
if database_scheme == "postgresql":
event.listen(
tool_database.engine,
"before_cursor_execute",
translate_postgresql_distinct_on,
retval=True,
)
try:
providers = ToolManager.list_default_builtin_providers(tenant_id)
finally:
if database_scheme == "postgresql":
event.remove(
tool_database.engine,
"before_cursor_execute",
translate_postgresql_distinct_on,
)
assert [provider.id for provider in providers] == [default_provider.id]
if database_scheme == "postgresql":
assert postgresql_statements
def test_list_providers_from_api_covers_builtin_api_workflow_and_mcp(monkeypatch: pytest.MonkeyPatch):
def test_list_providers_from_api_covers_builtin_api_workflow_and_mcp(
monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase
):
tenant_id = "00000000-0000-0000-0000-000000000001"
hardcoded_controller = SimpleNamespace(entity=SimpleNamespace(identity=SimpleNamespace(name="hardcoded")))
plugin_controller = object.__new__(PluginToolProviderController)
plugin_controller.entity = SimpleNamespace(identity=SimpleNamespace(name="plugin-provider"))
api_db_provider_good = SimpleNamespace(id="api-1")
api_db_provider_bad = SimpleNamespace(id="api-2")
api_controller = SimpleNamespace(provider_id="api-1")
api_db_provider_good = _api_provider(
provider_id="00000000-0000-0000-0000-000000000002", tenant_id=tenant_id, name="api-good"
)
api_db_provider_bad = _api_provider(
provider_id="00000000-0000-0000-0000-000000000003", tenant_id=tenant_id, name="api-bad"
)
api_controller = SimpleNamespace(provider_id=api_db_provider_good.id)
workflow_db_provider_good = SimpleNamespace(id="wf-1")
workflow_db_provider_bad = SimpleNamespace(id="wf-2")
workflow_controller = SimpleNamespace(provider_id="wf-1")
session = Mock()
session.scalars.side_effect = [
SimpleNamespace(all=lambda: [api_db_provider_good, api_db_provider_bad]),
SimpleNamespace(all=lambda: [workflow_db_provider_good, workflow_db_provider_bad]),
]
workflow_db_provider_good = _workflow_provider(
provider_id="00000000-0000-0000-0000-000000000004", tenant_id=tenant_id, name="workflow-good"
)
workflow_db_provider_bad = _workflow_provider(
provider_id="00000000-0000-0000-0000-000000000005", tenant_id=tenant_id, name="workflow-bad"
)
workflow_controller = SimpleNamespace(provider_id=workflow_db_provider_good.id)
tool_database.session.add_all(
[
api_db_provider_good,
api_db_provider_bad,
_api_provider(
provider_id="00000000-0000-0000-0000-000000000006",
tenant_id="00000000-0000-0000-0000-000000000099",
name="foreign-api",
),
workflow_db_provider_good,
workflow_db_provider_bad,
_workflow_provider(
provider_id="00000000-0000-0000-0000-000000000007",
tenant_id="00000000-0000-0000-0000-000000000099",
name="foreign-workflow",
),
]
)
tool_database.session.commit()
_setup_list_providers_from_api_mocks(
monkeypatch,
session=session,
tool_database=tool_database,
hardcoded_controller=hardcoded_controller,
plugin_controller=plugin_controller,
api_controller=api_controller,
workflow_controller=workflow_controller,
)
providers = ToolManager.list_providers_from_api(user_id="user-1", tenant_id="tenant-1", typ="")
providers = ToolManager.list_providers_from_api(user_id="user-1", tenant_id=tenant_id, typ="")
names = {provider.name for provider in providers}
assert {"hardcoded", "plugin-provider", "api-provider", "workflow-provider", "mcp-provider"} <= names
def test_get_api_provider_controller_returns_controller_and_credentials():
provider = SimpleNamespace(
id="api-1",
tenant_id="tenant-1",
name="api-provider",
description="desc",
credentials={"auth_type": "api_key_query"},
credentials_str='{"auth_type": "api_key_query", "api_key_value": "secret"}',
schema_type="openapi",
schema="schema",
tools=[],
icon='{"background": "#000", "content": "A"}',
privacy_policy="privacy",
custom_disclaimer="disclaimer",
)
def test_get_api_provider_controller_returns_controller_and_credentials(
monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase
):
tenant_id = "00000000-0000-0000-0000-000000000001"
provider = _api_provider(provider_id="00000000-0000-0000-0000-000000000002", tenant_id=tenant_id)
tool_database.session.add(provider)
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
controller = Mock()
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.session.scalar.return_value = provider
with patch(
"core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller
) as mock_from_db:
built_controller, credentials = ToolManager.get_api_provider_controller("tenant-1", "api-1")
with patch("core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller) as mock_from_db:
built_controller, credentials = ToolManager.get_api_provider_controller(tenant_id, provider.id)
assert built_controller is controller
assert credentials == provider.credentials
@ -711,83 +886,74 @@ def test_get_api_provider_controller_returns_controller_and_credentials():
controller.load_bundled_tools.assert_called_once_with(provider.tools)
def test_user_get_api_provider_masks_credentials_and_adds_labels():
provider = SimpleNamespace(
id="api-1",
tenant_id="tenant-1",
name="api-provider",
description="desc",
credentials={"auth_type": "api_key_query"},
credentials_str='{"auth_type": "api_key_query", "api_key_value": "secret"}',
schema_type="openapi",
schema="schema",
tools=[],
icon='{"background": "#000", "content": "A"}',
privacy_policy="privacy",
custom_disclaimer="disclaimer",
)
def test_user_get_api_provider_masks_credentials_and_adds_labels(
monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase
):
tenant_id = "00000000-0000-0000-0000-000000000001"
provider = _api_provider(provider_id="00000000-0000-0000-0000-000000000002", tenant_id=tenant_id)
tool_database.session.add(provider)
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
controller = Mock()
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.session.scalar.return_value = provider
with patch("core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller):
encrypter = Mock()
encrypter.decrypt.return_value = {"api_key_value": "secret"}
encrypter.mask_plugin_credentials.return_value = {"api_key_value": "***"}
with patch("core.tools.tool_manager.create_tool_provider_encrypter", return_value=(encrypter, Mock())):
with patch("core.tools.tool_manager.ToolLabelManager.get_tool_labels", return_value=["search"]):
user_payload = ToolManager.user_get_api_provider("api-provider", "tenant-1")
with patch("core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller):
encrypter = Mock()
encrypter.decrypt.return_value = {"api_key_value": "secret"}
encrypter.mask_plugin_credentials.return_value = {"api_key_value": "***"}
with patch("core.tools.tool_manager.create_tool_provider_encrypter", return_value=(encrypter, Mock())):
with patch("core.tools.tool_manager.ToolLabelManager.get_tool_labels", return_value=["search"]):
user_payload = ToolManager.user_get_api_provider(provider.name, tenant_id)
assert user_payload["credentials"]["api_key_value"] == "***"
assert user_payload["labels"] == ["search"]
def test_get_api_provider_controller_not_found_raises():
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.session.scalar.return_value = None
with pytest.raises(ToolProviderNotFoundError, match="api provider missing not found"):
ToolManager.get_api_provider_controller("tenant-1", "missing")
def test_get_api_provider_controller_not_found_raises(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase):
provider_id = "00000000-0000-0000-0000-000000000002"
tool_database.session.add(
_api_provider(
provider_id=provider_id,
tenant_id="00000000-0000-0000-0000-000000000099",
)
)
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
with pytest.raises(ToolProviderNotFoundError, match=f"api provider {provider_id} not found"):
ToolManager.get_api_provider_controller("00000000-0000-0000-0000-000000000001", provider_id)
def test_get_mcp_provider_controller_returns_controller():
def test_get_mcp_provider_controller_returns_controller(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase):
provider_entity = SimpleNamespace(provider_icon={"background": "#111", "content": "M"})
controller = Mock()
session = Mock()
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.engine = object()
with patch("core.tools.tool_manager.Session", return_value=_cm(session)):
with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls:
mock_service = mock_service_cls.return_value
mock_service.get_provider.return_value = provider_entity
with patch("core.tools.tool_manager.MCPToolProviderController.from_db", return_value=controller):
built = ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1")
assert built is controller
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls:
mock_service = mock_service_cls.return_value
mock_service.get_provider.return_value = provider_entity
with patch("core.tools.tool_manager.MCPToolProviderController.from_db", return_value=controller):
built = ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1")
assert built is controller
assert isinstance(mock_service_cls.call_args.kwargs["session"], Session)
def test_generate_mcp_tool_icon_url_returns_provider_icon():
def test_generate_mcp_tool_icon_url_returns_provider_icon(
monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase
):
provider_entity = SimpleNamespace(provider_icon={"background": "#111", "content": "M"})
session = Mock()
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.engine = object()
with patch("core.tools.tool_manager.Session", return_value=_cm(session)):
with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls:
mock_service = mock_service_cls.return_value
mock_service.get_provider_entity.return_value = provider_entity
assert ToolManager.generate_mcp_tool_icon_url("tenant-1", "mcp-1") == provider_entity.provider_icon
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls:
mock_service = mock_service_cls.return_value
mock_service.get_provider_entity.return_value = provider_entity
assert ToolManager.generate_mcp_tool_icon_url("tenant-1", "mcp-1") == provider_entity.provider_icon
assert isinstance(mock_service_cls.call_args.kwargs["session"], Session)
def test_get_mcp_provider_controller_missing_raises():
session = Mock()
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.engine = object()
with patch("core.tools.tool_manager.Session", return_value=_cm(session)):
with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls:
mock_service_cls.return_value.get_provider.side_effect = ValueError("missing")
with pytest.raises(ToolProviderNotFoundError, match="mcp provider mcp-1 not found"):
ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1")
def test_get_mcp_provider_controller_missing_raises(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase):
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls:
mock_service_cls.return_value.get_provider.side_effect = ValueError("missing")
with pytest.raises(ToolProviderNotFoundError, match="mcp provider mcp-1 not found"):
ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1")
def test_generate_tool_icon_urls_for_builtin_and_plugin():
@ -799,39 +965,49 @@ def test_generate_tool_icon_urls_for_builtin_and_plugin():
assert "/plugin/icon" in plugin_url
def test_generate_tool_icon_urls_for_workflow_and_api():
workflow_provider = SimpleNamespace(icon='{"background": "#222", "content": "W"}')
api_provider = SimpleNamespace(icon='{"background": "#333", "content": "A"}')
mock_engine = object()
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.engine = mock_engine
with patch("core.tools.tool_manager.Session") as mock_session_cls:
mock_session = MagicMock()
mock_session.scalar.side_effect = [workflow_provider, api_provider]
mock_session_cls.return_value.__enter__ = MagicMock(return_value=mock_session)
mock_session_cls.return_value.__exit__ = MagicMock(return_value=False)
assert ToolManager.generate_workflow_tool_icon_url("tenant-1", "wf-1") == {
"background": "#222",
"content": "W",
}
assert ToolManager.generate_api_tool_icon_url("tenant-1", "api-1") == {"background": "#333", "content": "A"}
# Verify sessions are created with the engine
assert mock_session_cls.call_count == 2
mock_session_cls.assert_called_with(mock_engine, expire_on_commit=False)
def test_generate_tool_icon_urls_for_workflow_and_api(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase):
tenant_id = "00000000-0000-0000-0000-000000000001"
workflow_provider = _workflow_provider(
provider_id="00000000-0000-0000-0000-000000000002",
tenant_id=tenant_id,
)
api_provider = _api_provider(
provider_id="00000000-0000-0000-0000-000000000003",
tenant_id=tenant_id,
icon='{"background":"#333","content":"A"}',
)
tool_database.session.add_all([workflow_provider, api_provider])
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
assert ToolManager.generate_workflow_tool_icon_url(tenant_id, workflow_provider.id) == {
"background": "#222",
"content": "W",
}
assert ToolManager.generate_api_tool_icon_url(tenant_id, api_provider.id) == {
"background": "#333",
"content": "A",
}
def test_generate_tool_icon_urls_missing_workflow_and_api_use_default():
mock_engine = object()
with patch("core.tools.tool_manager.db") as mock_db:
mock_db.engine = mock_engine
with patch("core.tools.tool_manager.Session") as mock_session_cls:
mock_session = MagicMock()
mock_session.scalar.return_value = None
mock_session_cls.return_value.__enter__ = MagicMock(return_value=mock_session)
mock_session_cls.return_value.__exit__ = MagicMock(return_value=False)
assert ToolManager.generate_workflow_tool_icon_url("tenant-1", "missing")["background"] == "#252525"
assert ToolManager.generate_api_tool_icon_url("tenant-1", "missing")["background"] == "#252525"
assert mock_session_cls.call_count == 2
def test_generate_tool_icon_urls_missing_workflow_and_api_use_default(
monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase
):
tenant_id = "00000000-0000-0000-0000-000000000001"
foreign_tenant_id = "00000000-0000-0000-0000-000000000099"
workflow_provider_id = "00000000-0000-0000-0000-000000000002"
api_provider_id = "00000000-0000-0000-0000-000000000003"
tool_database.session.add_all(
[
_workflow_provider(provider_id=workflow_provider_id, tenant_id=foreign_tenant_id),
_api_provider(provider_id=api_provider_id, tenant_id=foreign_tenant_id),
]
)
tool_database.session.commit()
monkeypatch.setattr("core.tools.tool_manager.db", tool_database)
assert ToolManager.generate_workflow_tool_icon_url(tenant_id, workflow_provider_id)["background"] == "#252525"
assert ToolManager.generate_api_tool_icon_url(tenant_id, api_provider_id)["background"] == "#252525"
def test_get_tool_icon_for_builtin_provider_variants():