mirror of
https://github.com/langgenius/dify.git
synced 2026-07-30 16:59:35 +08:00
test: use SQLite sessions in services tools (#39063)
This commit is contained in:
parent
a21d2b1201
commit
bea82fb57f
@ -1,82 +1,133 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
"""Unit tests for built-in tool management and its persisted credential state."""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.plugin.entities.plugin_daemon import CredentialType
|
||||
from models.tools import BuiltinToolProvider, ToolOAuthSystemClient, ToolOAuthTenantClient
|
||||
from services.tools import builtin_tools_manage_service as service_module
|
||||
from services.tools.builtin_tools_manage_service import BuiltinToolManageService
|
||||
|
||||
MODULE = "services.tools.builtin_tools_manage_service"
|
||||
|
||||
@pytest.fixture
|
||||
def repository_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session:
|
||||
"""Bind service-owned sessions to the shared SQLite session's engine."""
|
||||
monkeypatch.setattr(service_module, "db", SimpleNamespace(engine=sqlite_session.get_bind()))
|
||||
return sqlite_session
|
||||
|
||||
|
||||
def _mock_session(mock_session_cls):
|
||||
"""Helper: set up a Session context manager mock and return the inner session."""
|
||||
session = MagicMock()
|
||||
mock_session_cls.return_value.__enter__ = MagicMock(return_value=session)
|
||||
mock_session_cls.return_value.__exit__ = MagicMock(return_value=False)
|
||||
return session
|
||||
def _persist_provider(
|
||||
session: Session,
|
||||
*,
|
||||
credential_id: str = "cred-1",
|
||||
tenant_id: str = "tenant-1",
|
||||
user_id: str = "user-1",
|
||||
provider: str = "google",
|
||||
name: str = "Google 1",
|
||||
credentials: dict[str, str] | None = None,
|
||||
is_default: bool = False,
|
||||
) -> BuiltinToolProvider:
|
||||
db_provider = BuiltinToolProvider(
|
||||
tenant_id=tenant_id,
|
||||
user_id=user_id,
|
||||
provider=provider,
|
||||
name=name,
|
||||
encrypted_credentials=json.dumps(credentials or {"key": "encrypted"}),
|
||||
credential_type=CredentialType.API_KEY,
|
||||
is_default=is_default,
|
||||
)
|
||||
db_provider.id = credential_id
|
||||
session.add(db_provider)
|
||||
session.commit()
|
||||
return db_provider
|
||||
|
||||
|
||||
def _mock_sessionmaker(mock_sm_cls):
|
||||
"""Helper: set up a sessionmaker().begin() context manager mock and return the inner session."""
|
||||
session = MagicMock()
|
||||
mock_sm_cls.return_value.begin.return_value.__enter__ = MagicMock(return_value=session)
|
||||
mock_sm_cls.return_value.begin.return_value.__exit__ = MagicMock(return_value=False)
|
||||
return session
|
||||
def _persist_tenant_oauth_client(
|
||||
session: Session,
|
||||
*,
|
||||
tenant_id: str = "tenant-1",
|
||||
plugin_id: str = "langgenius/google",
|
||||
provider: str = "google",
|
||||
enabled: bool = True,
|
||||
encrypted_params: str = '{"encrypted": "data"}',
|
||||
) -> ToolOAuthTenantClient:
|
||||
client = ToolOAuthTenantClient(tenant_id=tenant_id, plugin_id=plugin_id, provider=provider)
|
||||
client.enabled = enabled
|
||||
client.encrypted_oauth_params = encrypted_params
|
||||
session.add(client)
|
||||
session.commit()
|
||||
return client
|
||||
|
||||
|
||||
def _persist_system_oauth_client(
|
||||
session: Session,
|
||||
*,
|
||||
plugin_id: str = "langgenius/google",
|
||||
provider: str = "google",
|
||||
) -> ToolOAuthSystemClient:
|
||||
client = ToolOAuthSystemClient(plugin_id=plugin_id, provider=provider, encrypted_oauth_params="enc")
|
||||
session.add(client)
|
||||
session.commit()
|
||||
return client
|
||||
|
||||
|
||||
class TestDeleteCustomOauthClientParams:
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_deletes_and_returns_success(self, mock_db, mock_sm_cls):
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
def test_deletes_matching_tenant_only(self, repository_session: Session) -> None:
|
||||
_persist_tenant_oauth_client(repository_session, tenant_id="tenant-1")
|
||||
_persist_tenant_oauth_client(repository_session, tenant_id="tenant-2")
|
||||
|
||||
result = BuiltinToolManageService.delete_custom_oauth_client_params("tenant-1", "google")
|
||||
|
||||
assert result == {"result": "success"}
|
||||
session.execute.assert_called_once()
|
||||
repository_session.expire_all()
|
||||
clients = repository_session.scalars(select(ToolOAuthTenantClient)).all()
|
||||
assert [client.tenant_id for client in clients] == ["tenant-2"]
|
||||
|
||||
|
||||
class TestListBuiltinToolProviderTools:
|
||||
@patch(f"{MODULE}.ToolLabelManager")
|
||||
@patch(f"{MODULE}.ToolTransformService")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_transforms_each_tool(self, mock_manager, mock_transform, mock_labels):
|
||||
mock_controller = MagicMock()
|
||||
mock_controller.get_tools.return_value = [MagicMock(), MagicMock()]
|
||||
mock_manager.get_builtin_provider.return_value = mock_controller
|
||||
mock_transform.convert_tool_entity_to_api_entity.return_value = MagicMock()
|
||||
def test_transforms_each_tool(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
controller = MagicMock()
|
||||
controller.get_tools.return_value = [MagicMock(), MagicMock()]
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
convert = MagicMock(return_value=MagicMock())
|
||||
monkeypatch.setattr(service_module.ToolTransformService, "convert_tool_entity_to_api_entity", convert)
|
||||
monkeypatch.setattr(service_module.ToolLabelManager, "get_tool_labels", MagicMock(return_value=[]))
|
||||
|
||||
result = BuiltinToolManageService.list_builtin_tool_provider_tools("tenant-1", "google")
|
||||
|
||||
assert len(result) == 2
|
||||
assert convert.call_count == 2
|
||||
|
||||
@patch(f"{MODULE}.ToolLabelManager")
|
||||
@patch(f"{MODULE}.ToolTransformService")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_empty_tools(self, mock_manager, mock_transform, mock_labels):
|
||||
mock_controller = MagicMock()
|
||||
mock_controller.get_tools.return_value = []
|
||||
mock_manager.get_builtin_provider.return_value = mock_controller
|
||||
def test_empty_tools(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
controller = MagicMock()
|
||||
controller.get_tools.return_value = list[object]()
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
|
||||
assert BuiltinToolManageService.list_builtin_tool_provider_tools("t", "p") == []
|
||||
|
||||
|
||||
class TestGetBuiltinToolProviderInfo:
|
||||
@patch(f"{MODULE}.ToolTransformService")
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.get_builtin_provider")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_raises_when_not_found(self, mock_manager, mock_get, mock_transform):
|
||||
mock_get.return_value = None
|
||||
def test_raises_when_not_found(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(BuiltinToolManageService, "get_builtin_provider", MagicMock(return_value=None))
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=MagicMock()))
|
||||
|
||||
with pytest.raises(ValueError, match="you have not added provider"):
|
||||
BuiltinToolManageService.get_builtin_tool_provider_info("t", "no")
|
||||
|
||||
@patch(f"{MODULE}.ToolTransformService")
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.get_builtin_provider")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_clears_original_credentials(self, mock_manager, mock_get, mock_transform):
|
||||
mock_get.return_value = MagicMock()
|
||||
def test_clears_original_credentials(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(BuiltinToolManageService, "get_builtin_provider", MagicMock(return_value=MagicMock()))
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=MagicMock()))
|
||||
entity = MagicMock()
|
||||
mock_transform.builtin_provider_to_user_provider.return_value = entity
|
||||
monkeypatch.setattr(
|
||||
service_module.ToolTransformService,
|
||||
"builtin_provider_to_user_provider",
|
||||
MagicMock(return_value=entity),
|
||||
)
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_info("t", "google")
|
||||
|
||||
@ -84,21 +135,26 @@ class TestGetBuiltinToolProviderInfo:
|
||||
|
||||
|
||||
class TestListBuiltinProviderCredentialsSchema:
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_returns_schema(self, mock_manager):
|
||||
mock_manager.get_builtin_provider.return_value.get_credentials_schema_by_type.return_value = [{"f": "k"}]
|
||||
def test_returns_schema(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
controller = MagicMock()
|
||||
controller.get_credentials_schema_by_type.return_value = [{"f": "k"}]
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
|
||||
result = BuiltinToolManageService.list_builtin_provider_credentials_schema("g", "api_key", "t")
|
||||
result = BuiltinToolManageService.list_builtin_provider_credentials_schema("g", CredentialType.API_KEY, "t")
|
||||
|
||||
assert result == [{"f": "k"}]
|
||||
|
||||
|
||||
class TestGetBuiltinToolProviderIcon:
|
||||
@patch(f"{MODULE}.Path")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_returns_bytes_and_mime(self, mock_manager, mock_path):
|
||||
mock_manager.get_hardcoded_provider_icon.return_value = ("/icon.svg", "image/svg+xml")
|
||||
mock_path.return_value.read_bytes.return_value = b"<svg/>"
|
||||
def test_returns_bytes_and_mime(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(
|
||||
service_module.ToolManager,
|
||||
"get_hardcoded_provider_icon",
|
||||
MagicMock(return_value=("/icon.svg", "image/svg+xml")),
|
||||
)
|
||||
path = MagicMock()
|
||||
path.return_value.read_bytes.return_value = b"<svg/>"
|
||||
monkeypatch.setattr(service_module, "Path", path)
|
||||
|
||||
icon, mime = BuiltinToolManageService.get_builtin_tool_provider_icon("google")
|
||||
|
||||
@ -107,164 +163,140 @@ class TestGetBuiltinToolProviderIcon:
|
||||
|
||||
|
||||
class TestIsOauthSystemClientExists:
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_true_when_exists(self, mock_db, mock_session_cls):
|
||||
session = _mock_session(mock_session_cls)
|
||||
session.scalar.return_value = MagicMock()
|
||||
def test_true_when_exists(self, repository_session: Session) -> None:
|
||||
_persist_system_oauth_client(repository_session)
|
||||
|
||||
assert BuiltinToolManageService.is_oauth_system_client_exists("google") is True
|
||||
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_false_when_missing(self, mock_db, mock_session_cls):
|
||||
session = _mock_session(mock_session_cls)
|
||||
session.scalar.return_value = None
|
||||
def test_false_when_missing(self, repository_session: Session) -> None:
|
||||
_persist_system_oauth_client(repository_session, plugin_id="langgenius/slack", provider="slack")
|
||||
|
||||
assert BuiltinToolManageService.is_oauth_system_client_exists("google") is False
|
||||
|
||||
|
||||
class TestIsOauthCustomClientEnabled:
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_true_when_enabled(self, mock_db, mock_session_cls):
|
||||
session = _mock_session(mock_session_cls)
|
||||
session.scalar.return_value = MagicMock(enabled=True)
|
||||
def test_true_when_enabled(self, repository_session: Session) -> None:
|
||||
_persist_tenant_oauth_client(repository_session)
|
||||
|
||||
assert BuiltinToolManageService.is_oauth_custom_client_enabled("t", "g") is True
|
||||
assert BuiltinToolManageService.is_oauth_custom_client_enabled("tenant-1", "google") is True
|
||||
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_false_when_none(self, mock_db, mock_session_cls):
|
||||
session = _mock_session(mock_session_cls)
|
||||
session.scalar.return_value = None
|
||||
def test_false_when_disabled_or_other_tenant(self, repository_session: Session) -> None:
|
||||
_persist_tenant_oauth_client(repository_session, tenant_id="tenant-1", enabled=False)
|
||||
_persist_tenant_oauth_client(repository_session, tenant_id="tenant-2", enabled=True)
|
||||
|
||||
assert BuiltinToolManageService.is_oauth_custom_client_enabled("t", "g") is False
|
||||
assert BuiltinToolManageService.is_oauth_custom_client_enabled("tenant-1", "google") is False
|
||||
|
||||
|
||||
class TestDeleteBuiltinToolProvider:
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_raises_when_not_found(self, mock_db, mock_sm_cls, mock_tm, mock_enc):
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
session.scalar.return_value = None
|
||||
|
||||
def test_raises_when_not_found(self, repository_session: Session) -> None:
|
||||
with pytest.raises(ValueError, match="you have not added provider"):
|
||||
BuiltinToolManageService.delete_builtin_tool_provider("t", "p", "id")
|
||||
BuiltinToolManageService.delete_builtin_tool_provider("tenant-1", "google", "missing")
|
||||
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_deletes_provider_and_clears_cache(self, mock_db, mock_sm_cls, mock_tm, mock_enc):
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
db_provider = MagicMock()
|
||||
session.scalar.return_value = db_provider
|
||||
mock_cache = MagicMock()
|
||||
mock_enc.return_value = (MagicMock(), mock_cache)
|
||||
def test_deletes_provider_and_clears_cache(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
repository_session: Session,
|
||||
) -> None:
|
||||
_persist_provider(repository_session, credential_id="cred-1")
|
||||
_persist_provider(repository_session, credential_id="cred-other", tenant_id="tenant-2")
|
||||
cache = MagicMock()
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=MagicMock()))
|
||||
monkeypatch.setattr(
|
||||
BuiltinToolManageService,
|
||||
"create_tool_encrypter",
|
||||
MagicMock(return_value=(MagicMock(), cache)),
|
||||
)
|
||||
|
||||
result = BuiltinToolManageService.delete_builtin_tool_provider("t", "p", "c")
|
||||
result = BuiltinToolManageService.delete_builtin_tool_provider("tenant-1", "google", "cred-1")
|
||||
|
||||
assert result == {"result": "success"}
|
||||
session.delete.assert_called_once_with(db_provider)
|
||||
mock_cache.delete.assert_called_once()
|
||||
cache.delete.assert_called_once()
|
||||
repository_session.expire_all()
|
||||
assert repository_session.get(BuiltinToolProvider, "cred-1") is None
|
||||
assert repository_session.get(BuiltinToolProvider, "cred-other") is not None
|
||||
|
||||
|
||||
class TestSetDefaultProvider:
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_raises_when_not_found(self, mock_db, mock_sm_cls):
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
session.scalar.return_value = None
|
||||
|
||||
def test_raises_when_not_found(self, repository_session: Session) -> None:
|
||||
with pytest.raises(ValueError, match="provider not found"):
|
||||
BuiltinToolManageService.set_default_provider("t", "p", "id")
|
||||
BuiltinToolManageService.set_default_provider("tenant-1", "google", "missing")
|
||||
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_sets_default_and_clears_old(self, mock_db, mock_sm_cls):
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
target = MagicMock()
|
||||
session.scalar.return_value = target
|
||||
def test_sets_target_and_clears_only_same_tenant_defaults(self, repository_session: Session) -> None:
|
||||
_persist_provider(
|
||||
repository_session,
|
||||
credential_id="target",
|
||||
user_id="user-2",
|
||||
name="Google target",
|
||||
)
|
||||
_persist_provider(repository_session, credential_id="old", name="Google old", is_default=True)
|
||||
_persist_provider(
|
||||
repository_session,
|
||||
credential_id="other-tenant",
|
||||
tenant_id="tenant-2",
|
||||
name="Other tenant",
|
||||
is_default=True,
|
||||
)
|
||||
|
||||
result = BuiltinToolManageService.set_default_provider("t", "p", "id")
|
||||
result = BuiltinToolManageService.set_default_provider("tenant-1", "google", "target")
|
||||
|
||||
assert result == {"result": "success"}
|
||||
repository_session.expire_all()
|
||||
target = repository_session.get(BuiltinToolProvider, "target")
|
||||
old = repository_session.get(BuiltinToolProvider, "old")
|
||||
other_tenant = repository_session.get(BuiltinToolProvider, "other-tenant")
|
||||
assert target is not None
|
||||
assert target.is_default is True
|
||||
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_clear_default_is_tenant_scoped_not_user_scoped(self, mock_db, mock_sm_cls):
|
||||
# Regression: clearing prior defaults must NOT filter by user_id, otherwise
|
||||
# two workspace members can each leave their own credential as default at
|
||||
# the same time (the default flag is tenant-scoped, not per-user).
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
session.scalar.return_value = MagicMock()
|
||||
|
||||
BuiltinToolManageService.set_default_provider("tenant-1", "google", "cred-id")
|
||||
|
||||
session.execute.assert_called_once()
|
||||
update_stmt = session.execute.call_args.args[0]
|
||||
compiled = str(update_stmt.compile(compile_kwargs={"literal_binds": True}))
|
||||
assert "user_id" not in compiled
|
||||
assert "tenant_id" in compiled
|
||||
assert "provider" in compiled
|
||||
assert old is not None
|
||||
assert old.is_default is False
|
||||
assert other_tenant is not None
|
||||
assert other_tenant.is_default is True
|
||||
|
||||
|
||||
class TestUpdateBuiltinToolProvider:
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_raises_when_provider_not_exists(self, mock_db, mock_sm_cls):
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
session.scalar.return_value = None
|
||||
|
||||
def test_raises_when_provider_not_exists(self, repository_session: Session) -> None:
|
||||
with pytest.raises(ValueError, match="you have not added provider"):
|
||||
BuiltinToolManageService.update_builtin_tool_provider("u", "t", "p", "c")
|
||||
BuiltinToolManageService.update_builtin_tool_provider("u", "tenant-1", "google", "missing")
|
||||
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter")
|
||||
@patch(f"{MODULE}.CredentialType")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
@patch(f"{MODULE}.sessionmaker")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_updates_credentials_and_commits(self, mock_db, mock_sm_cls, mock_tm, mock_cred_type, mock_enc):
|
||||
session = _mock_sessionmaker(mock_sm_cls)
|
||||
db_provider = MagicMock(credential_type="api_key", credentials="{}")
|
||||
session.scalar.return_value = db_provider
|
||||
def test_updates_persisted_credentials_and_clears_cache(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
repository_session: Session,
|
||||
) -> None:
|
||||
_persist_provider(repository_session, credentials={"key": "old"})
|
||||
controller = MagicMock(need_credentials=True)
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
encrypter = MagicMock()
|
||||
encrypter.decrypt.return_value = {"key": "old"}
|
||||
encrypter.encrypt.return_value = {"key": "new"}
|
||||
cache = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
BuiltinToolManageService,
|
||||
"create_tool_encrypter",
|
||||
MagicMock(return_value=(encrypter, cache)),
|
||||
)
|
||||
|
||||
mock_cred_instance = MagicMock()
|
||||
mock_cred_instance.is_editable.return_value = True
|
||||
mock_cred_instance.is_validate_allowed.return_value = False
|
||||
mock_cred_type.of.return_value = mock_cred_instance
|
||||
|
||||
mock_controller = MagicMock(need_credentials=True)
|
||||
mock_tm.get_builtin_provider.return_value = mock_controller
|
||||
|
||||
mock_encrypter = MagicMock()
|
||||
mock_encrypter.decrypt.return_value = {"key": "old"}
|
||||
mock_encrypter.encrypt.return_value = {"key": "new"}
|
||||
mock_cache = MagicMock()
|
||||
mock_enc.return_value = (mock_encrypter, mock_cache)
|
||||
|
||||
result = BuiltinToolManageService.update_builtin_tool_provider("u", "t", "p", "c", credentials={"key": "val"})
|
||||
result = BuiltinToolManageService.update_builtin_tool_provider(
|
||||
"u", "tenant-1", "google", "cred-1", credentials={"key": "value"}
|
||||
)
|
||||
|
||||
assert result == {"result": "success"}
|
||||
mock_cache.delete.assert_called_once()
|
||||
controller.validate_credentials.assert_called_once_with("u", {"key": "value"})
|
||||
cache.delete.assert_called_once()
|
||||
repository_session.expire_all()
|
||||
provider = repository_session.get(BuiltinToolProvider, "cred-1")
|
||||
assert provider is not None
|
||||
assert provider.credentials == {"key": "new"}
|
||||
|
||||
|
||||
class TestGetOauthClientSchema:
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.get_custom_oauth_client_params", return_value={})
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.is_oauth_system_client_exists", return_value=False)
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.is_oauth_custom_client_enabled", return_value=True)
|
||||
@patch(f"{MODULE}.dify_config")
|
||||
@patch(f"{MODULE}.PluginService")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_returns_schema_dict(self, mock_tm, mock_plugin, mock_config, mock_enabled, mock_sys, mock_params):
|
||||
mock_config.CONSOLE_API_URL = "https://api.example.com"
|
||||
mock_controller = MagicMock()
|
||||
mock_controller.get_oauth_client_schema.return_value = []
|
||||
mock_tm.get_builtin_provider.return_value = mock_controller
|
||||
def test_returns_schema_dict(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
controller = MagicMock()
|
||||
controller.get_oauth_client_schema.return_value = list[object]()
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
monkeypatch.setattr(BuiltinToolManageService, "is_oauth_custom_client_enabled", MagicMock(return_value=True))
|
||||
monkeypatch.setattr(BuiltinToolManageService, "is_oauth_system_client_exists", MagicMock(return_value=False))
|
||||
monkeypatch.setattr(BuiltinToolManageService, "get_custom_oauth_client_params", MagicMock(return_value={}))
|
||||
monkeypatch.setattr(service_module.dify_config, "CONSOLE_API_URL", "https://api.example.com")
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_oauth_client_schema("t", "google")
|
||||
|
||||
@ -274,87 +306,91 @@ class TestGetOauthClientSchema:
|
||||
|
||||
|
||||
class TestGetOauthClient:
|
||||
@patch(f"{MODULE}.PluginService")
|
||||
@patch(f"{MODULE}.create_provider_encrypter")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_returns_user_client_params_when_exists(
|
||||
self, mock_db, mock_session_cls, mock_tm, mock_create_enc, mock_plugin
|
||||
):
|
||||
session = _mock_session(mock_session_cls)
|
||||
mock_controller = MagicMock()
|
||||
mock_controller.get_oauth_client_schema.return_value = []
|
||||
mock_tm.get_builtin_provider.return_value = mock_controller
|
||||
def test_returns_tenant_client_params_when_exists(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
repository_session: Session,
|
||||
) -> None:
|
||||
_persist_tenant_oauth_client(repository_session)
|
||||
controller = MagicMock()
|
||||
controller.get_oauth_client_schema.return_value = list[object]()
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
encrypter = MagicMock()
|
||||
encrypter.decrypt.return_value = {"client_id": "id", "client_secret": "secret"}
|
||||
monkeypatch.setattr(
|
||||
service_module, "create_provider_encrypter", MagicMock(return_value=(encrypter, MagicMock()))
|
||||
)
|
||||
|
||||
mock_encrypter = MagicMock()
|
||||
mock_encrypter.decrypt.return_value = {"client_id": "id", "client_secret": "secret"}
|
||||
mock_create_enc.return_value = (mock_encrypter, MagicMock())
|
||||
|
||||
user_client = MagicMock(oauth_params='{"encrypted": "data"}')
|
||||
session.scalar.return_value = user_client
|
||||
|
||||
result = BuiltinToolManageService.get_oauth_client("t", "google")
|
||||
result = BuiltinToolManageService.get_oauth_client("tenant-1", "google")
|
||||
|
||||
assert result == {"client_id": "id", "client_secret": "secret"}
|
||||
encrypter.decrypt.assert_called_once_with({"encrypted": "data"})
|
||||
|
||||
@patch(f"{MODULE}.decrypt_system_params", return_value={"sys_key": "sys_val"})
|
||||
@patch(f"{MODULE}.PluginService")
|
||||
@patch(f"{MODULE}.create_provider_encrypter")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_falls_back_to_system_client(
|
||||
self, mock_db, mock_session_cls, mock_tm, mock_create_enc, mock_plugin, mock_decrypt
|
||||
):
|
||||
session = _mock_session(mock_session_cls)
|
||||
mock_controller = MagicMock()
|
||||
mock_controller.get_oauth_client_schema.return_value = []
|
||||
mock_tm.get_builtin_provider.return_value = mock_controller
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
repository_session: Session,
|
||||
) -> None:
|
||||
_persist_system_oauth_client(repository_session)
|
||||
controller = MagicMock()
|
||||
controller.get_oauth_client_schema.return_value = list[object]()
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
monkeypatch.setattr(
|
||||
service_module, "create_provider_encrypter", MagicMock(return_value=(MagicMock(), MagicMock()))
|
||||
)
|
||||
decrypt = MagicMock(return_value={"sys_key": "sys_val"})
|
||||
monkeypatch.setattr(service_module, "decrypt_system_params", decrypt)
|
||||
|
||||
mock_create_enc.return_value = (MagicMock(), MagicMock())
|
||||
|
||||
system_client = MagicMock(encrypted_oauth_params="enc")
|
||||
session.scalar.side_effect = [None, system_client]
|
||||
|
||||
result = BuiltinToolManageService.get_oauth_client("t", "google")
|
||||
result = BuiltinToolManageService.get_oauth_client("tenant-1", "google")
|
||||
|
||||
assert result == {"sys_key": "sys_val"}
|
||||
decrypt.assert_called_once_with("enc")
|
||||
|
||||
|
||||
class TestSaveCustomOauthClientParams:
|
||||
def test_returns_early_when_no_params(self):
|
||||
def test_returns_early_when_no_params(self) -> None:
|
||||
result = BuiltinToolManageService.save_custom_oauth_client_params("t", "p")
|
||||
assert result == {"result": "success"}
|
||||
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_raises_when_provider_not_found(self, mock_tm):
|
||||
mock_tm.get_builtin_provider.return_value = None
|
||||
def test_raises_when_provider_not_found(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=None))
|
||||
|
||||
with pytest.raises((ValueError, Exception), match="not found|Provider"):
|
||||
BuiltinToolManageService.save_custom_oauth_client_params("t", "p", enable_oauth_custom_client=True)
|
||||
|
||||
|
||||
class TestGetCustomOauthClientParams:
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_returns_empty_when_none(self, mock_db, mock_session_cls):
|
||||
session = _mock_session(mock_session_cls)
|
||||
session.scalar.return_value = None
|
||||
def test_returns_empty_when_none(self, repository_session: Session) -> None:
|
||||
_persist_tenant_oauth_client(repository_session, tenant_id="other-tenant")
|
||||
|
||||
result = BuiltinToolManageService.get_custom_oauth_client_params("t", "p")
|
||||
result = BuiltinToolManageService.get_custom_oauth_client_params("tenant-1", "google")
|
||||
|
||||
assert result == {}
|
||||
|
||||
|
||||
class TestGetBuiltinToolProviderCredentialInfo:
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.is_oauth_custom_client_enabled", return_value=False)
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.get_builtin_tool_provider_credentials", return_value=[])
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
def test_returns_credential_info(self, mock_tm, mock_creds, mock_oauth):
|
||||
mock_tm.get_builtin_provider.return_value.get_supported_credential_types.return_value = ["api-key"]
|
||||
def test_returns_credential_info(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
repository_session: Session,
|
||||
) -> None:
|
||||
controller = MagicMock()
|
||||
controller.get_supported_credential_types.return_value = ["api-key"]
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
monkeypatch.setattr(
|
||||
BuiltinToolManageService,
|
||||
"get_builtin_tool_provider_credentials",
|
||||
MagicMock(return_value=[]),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
BuiltinToolManageService,
|
||||
"is_oauth_custom_client_enabled",
|
||||
MagicMock(return_value=False),
|
||||
)
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_credential_info("t", "google", session=MagicMock())
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_credential_info(
|
||||
"tenant-1", "google", session=repository_session
|
||||
)
|
||||
|
||||
assert result.credentials == []
|
||||
assert result.supported_credential_types == ["api-key"]
|
||||
@ -362,113 +398,84 @@ class TestGetBuiltinToolProviderCredentialInfo:
|
||||
|
||||
|
||||
class TestGetBuiltinToolProviderCredentials:
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_returns_empty_when_no_providers(self, mock_db):
|
||||
mock_db.session.no_autoflush.__enter__ = MagicMock(return_value=None)
|
||||
mock_db.session.no_autoflush.__exit__ = MagicMock(return_value=False)
|
||||
mock_db.session.scalars.return_value.all.return_value = []
|
||||
def test_returns_empty_when_no_providers(self, repository_session: Session) -> None:
|
||||
_persist_provider(repository_session, credential_id="other", tenant_id="other-tenant")
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session)
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_credentials(
|
||||
"tenant-1", "google", session=repository_session
|
||||
)
|
||||
|
||||
assert result == []
|
||||
|
||||
@patch(f"{MODULE}.ToolTransformService")
|
||||
@patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter")
|
||||
@patch(f"{MODULE}.ToolManager")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_returns_credential_entities(self, mock_db, mock_tm, mock_enc, mock_transform):
|
||||
mock_db.session.no_autoflush.__enter__ = MagicMock(return_value=None)
|
||||
mock_db.session.no_autoflush.__exit__ = MagicMock(return_value=False)
|
||||
|
||||
provider = MagicMock(provider="google", is_default=False)
|
||||
mock_db.session.scalars.return_value.all.return_value = [provider]
|
||||
|
||||
mock_encrypter = MagicMock()
|
||||
mock_encrypter.decrypt.return_value = {"key": "decrypted"}
|
||||
mock_encrypter.mask_plugin_credentials.return_value = {"key": "***"}
|
||||
mock_enc.return_value = (mock_encrypter, MagicMock())
|
||||
|
||||
def test_returns_tenant_scoped_credential_entities(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
repository_session: Session,
|
||||
) -> None:
|
||||
_persist_provider(repository_session, is_default=False)
|
||||
_persist_provider(repository_session, credential_id="other", tenant_id="other-tenant")
|
||||
controller = MagicMock()
|
||||
monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller))
|
||||
encrypter = MagicMock()
|
||||
encrypter.decrypt.return_value = {"key": "decrypted"}
|
||||
encrypter.mask_plugin_credentials.return_value = {"key": "***"}
|
||||
monkeypatch.setattr(
|
||||
BuiltinToolManageService,
|
||||
"create_tool_encrypter",
|
||||
MagicMock(return_value=(encrypter, MagicMock())),
|
||||
)
|
||||
credential_entity = MagicMock()
|
||||
mock_transform.convert_builtin_provider_to_credential_entity.return_value = credential_entity
|
||||
convert = MagicMock(return_value=credential_entity)
|
||||
monkeypatch.setattr(
|
||||
service_module.ToolTransformService,
|
||||
"convert_builtin_provider_to_credential_entity",
|
||||
convert,
|
||||
)
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session)
|
||||
result = BuiltinToolManageService.get_builtin_tool_provider_credentials(
|
||||
"tenant-1", "google", session=repository_session
|
||||
)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0] is credential_entity
|
||||
assert provider.is_default is True
|
||||
assert result == [credential_entity]
|
||||
converted_provider = convert.call_args.kwargs["provider"]
|
||||
assert isinstance(converted_provider, BuiltinToolProvider)
|
||||
assert converted_provider.tenant_id == "tenant-1"
|
||||
assert converted_provider.is_default is True
|
||||
|
||||
|
||||
class TestGetBuiltinProvider:
|
||||
@patch(f"{MODULE}.ToolProviderID")
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_returns_none_when_not_found(self, mock_db, mock_session_cls, mock_prov_id):
|
||||
session = _mock_session(mock_session_cls)
|
||||
mock_prov_id.return_value.provider_name = "google"
|
||||
mock_prov_id.return_value.organization = "langgenius"
|
||||
session.scalar.return_value = None
|
||||
def test_returns_none_when_not_found(self, repository_session: Session) -> None:
|
||||
assert BuiltinToolManageService.get_builtin_provider("google", "tenant-1") is None
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_provider("google", "t")
|
||||
def test_returns_langgenius_provider_for_matching_tenant(self, repository_session: Session) -> None:
|
||||
_persist_provider(repository_session, tenant_id="tenant-1", provider="google")
|
||||
_persist_provider(repository_session, credential_id="other", tenant_id="tenant-2", provider="google")
|
||||
|
||||
assert result is None
|
||||
result = BuiltinToolManageService.get_builtin_provider("google", "tenant-1")
|
||||
|
||||
@patch(f"{MODULE}.ToolProviderID")
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_returns_provider_for_langgenius_org(self, mock_db, mock_session_cls, mock_prov_id):
|
||||
session = _mock_session(mock_session_cls)
|
||||
mock_prov_id.return_value.provider_name = "google"
|
||||
mock_prov_id.return_value.organization = "langgenius"
|
||||
db_provider = MagicMock(provider="google")
|
||||
mock_prov_id_result = MagicMock()
|
||||
mock_prov_id_result.to_string.return_value = "langgenius/google/google"
|
||||
assert result is not None
|
||||
assert result.id == "cred-1"
|
||||
assert result.provider == "langgenius/google/google"
|
||||
|
||||
def prov_id_side_effect(name):
|
||||
m = MagicMock()
|
||||
m.provider_name = "google"
|
||||
m.organization = "langgenius"
|
||||
m.to_string.return_value = "langgenius/google/google"
|
||||
m.plugin_id = "langgenius/google"
|
||||
return m
|
||||
def test_returns_non_langgenius_provider(self, repository_session: Session) -> None:
|
||||
full_provider = "third-party/custom/custom-tool"
|
||||
_persist_provider(repository_session, provider=full_provider)
|
||||
|
||||
mock_prov_id.side_effect = prov_id_side_effect
|
||||
session.scalar.return_value = db_provider
|
||||
result = BuiltinToolManageService.get_builtin_provider(full_provider, "tenant-1")
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_provider("google", "t")
|
||||
assert result is not None
|
||||
assert result.id == "cred-1"
|
||||
assert result.provider == full_provider
|
||||
|
||||
assert result is db_provider
|
||||
def test_falls_back_on_provider_id_parse_exception(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
repository_session: Session,
|
||||
) -> None:
|
||||
_persist_provider(repository_session, provider="old-provider")
|
||||
monkeypatch.setattr(service_module, "ToolProviderID", MagicMock(side_effect=Exception("parse error")))
|
||||
|
||||
@patch(f"{MODULE}.ToolProviderID")
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_returns_provider_for_non_langgenius_org(self, mock_db, mock_session_cls, mock_prov_id):
|
||||
session = _mock_session(mock_session_cls)
|
||||
result = BuiltinToolManageService.get_builtin_provider("old-provider", "tenant-1")
|
||||
|
||||
def prov_id_side_effect(name):
|
||||
m = MagicMock()
|
||||
m.provider_name = "custom-tool"
|
||||
m.organization = "third-party"
|
||||
m.to_string.return_value = "third-party/custom/custom-tool"
|
||||
m.plugin_id = "third-party/custom"
|
||||
return m
|
||||
|
||||
mock_prov_id.side_effect = prov_id_side_effect
|
||||
db_provider = MagicMock(provider="third-party/custom/custom-tool")
|
||||
session.scalar.return_value = db_provider
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_provider("third-party/custom/custom-tool", "t")
|
||||
|
||||
assert result is db_provider
|
||||
|
||||
@patch(f"{MODULE}.ToolProviderID")
|
||||
@patch(f"{MODULE}.Session")
|
||||
@patch(f"{MODULE}.db")
|
||||
def test_falls_back_on_exception(self, mock_db, mock_session_cls, mock_prov_id):
|
||||
session = _mock_session(mock_session_cls)
|
||||
mock_prov_id.side_effect = Exception("parse error")
|
||||
fallback = MagicMock()
|
||||
session.scalar.return_value = fallback
|
||||
|
||||
result = BuiltinToolManageService.get_builtin_provider("old-provider", "t")
|
||||
|
||||
assert result is fallback
|
||||
assert result is not None
|
||||
assert result.id == "cred-1"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user