mirror of
https://github.com/langgenius/dify.git
synced 2026-08-01 09:50:50 +08:00
test: use SQLite sessions in core tools (#39068)
This commit is contained in:
parent
9ee7750411
commit
bab3db7f77
@ -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():
|
||||
|
||||
Loading…
Reference in New Issue
Block a user