test: use sqlite3 session in test_tool_providers (#38754)

This commit is contained in:
Asuka Minato 2026-07-16 21:28:26 +09:00 committed by GitHub
parent 6ad1a66a50
commit 415c0db22e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -6,18 +6,22 @@ import builtins
import importlib
from contextlib import ExitStack, contextmanager
from inspect import unwrap
from types import ModuleType, SimpleNamespace
from unittest.mock import ANY, MagicMock, patch
from types import ModuleType
from unittest.mock import MagicMock, patch
from uuid import NAMESPACE_URL, uuid5
import pytest
from flask import Flask
from flask.views import MethodView
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from core.tools.entities.api_entities import ToolProviderApiEntity as CoreToolProviderApiEntity
from core.tools.entities.common_entities import I18nObject
from core.tools.entities.tool_entities import ToolParameter
from models import Account
from models import Account, BuiltinToolProvider, Tenant, TenantAccountJoin
from models.account import TenantAccountRole
from models.credential_permission import CredentialPermission
from models.enums import PermissionEnum
if not hasattr(builtins, "MethodView"):
builtins.MethodView = MethodView # type: ignore[attr-defined]
@ -27,13 +31,6 @@ _CONTROLLER_MODULE: ModuleType | None = None
_WRAPS_MODULE: ModuleType | None = None
@contextmanager
def _mock_db():
mock_session = SimpleNamespace(scalar=lambda *args, **kwargs: True)
with patch("extensions.ext_database.db.session", mock_session):
yield
@pytest.fixture
def app() -> Flask:
flask_app = Flask(__name__)
@ -69,8 +66,7 @@ def controller_module(monkeypatch: pytest.MonkeyPatch):
with ExitStack() as stack:
for target, value in patch_targets:
stack.enter_context(patch(target, value))
with _mock_db():
_CONTROLLER_MODULE = importlib.import_module(module_name)
_CONTROLLER_MODULE = importlib.import_module(module_name)
module = _CONTROLLER_MODULE
@ -88,11 +84,79 @@ def controller_module(monkeypatch: pytest.MonkeyPatch):
def _mock_account(user_id: str = "user-123") -> Account:
user = Account(name="Test User", email=f"{user_id}@example.com")
user.id = user_id
user.id = _stable_uuid(f"account:{user_id}")
user.role = TenantAccountRole.NORMAL
return user
def _stable_uuid(value: str) -> str:
return str(uuid5(NAMESPACE_URL, value))
def _persist_workspace(session: Session, user: Account, tenant_name: str) -> Tenant:
tenant = Tenant(name=tenant_name)
tenant.id = _stable_uuid(f"tenant:{tenant_name}")
membership = TenantAccountJoin(
tenant_id=tenant.id,
account_id=user.id,
current=True,
role=TenantAccountRole.NORMAL,
)
session.add_all([user, tenant, membership])
session.commit()
return tenant
def _provider_credential(
*,
tenant_id: str,
user_id: str,
credential_name: str,
visibility: PermissionEnum = PermissionEnum.ALL_TEAM,
) -> BuiltinToolProvider:
provider = BuiltinToolProvider(
name=credential_name,
tenant_id=tenant_id,
user_id=user_id,
provider="demo",
encrypted_credentials='{"api_key": "sk-secret"}',
visibility=visibility,
)
provider.id = _stable_uuid(f"credential:{tenant_id}:{credential_name}")
return provider
@contextmanager
def _bind_database_session(session: Session):
database_session = scoped_session(
sessionmaker(bind=session.get_bind(), expire_on_commit=False),
)
try:
with patch("extensions.ext_database.db.session", database_session):
yield
finally:
database_session.remove()
@contextmanager
def _mock_credential_encryption(controller_module: ModuleType):
encrypter = MagicMock()
encrypter.decrypt.side_effect = lambda credentials: credentials
encrypter.mask_plugin_credentials.return_value = {"api_key": "[__HIDDEN__]"}
with (
patch(
"services.tools.builtin_tools_manage_service.ToolManager.get_builtin_provider",
return_value=MagicMock(),
),
patch.object(
controller_module.BuiltinToolManageService,
"create_tool_encrypter",
return_value=(encrypter, MagicMock()),
),
):
yield
def _set_current_account(
monkeypatch: pytest.MonkeyPatch,
controller_module: ModuleType,
@ -190,23 +254,6 @@ def _provider_list_item(
return service_payload, provider.model_dump(mode="json", exclude_unset=True)
def _credential_response(controller_module: ModuleType, credential_id: str = "cred-1") -> tuple[dict, dict]:
expected = {
"id": credential_id,
"name": "Credential",
"provider": "demo",
"credential_type": controller_module.CredentialType.API_KEY,
"is_default": False,
"credentials": {},
"visibility": "all_team_members",
"created_by": "",
"partial_member_list": [],
"from_other_member": False,
}
credential = controller_module.ToolProviderCredentialApiEntity.model_validate(expected)
return credential.model_dump(mode="json"), credential.model_dump(mode="json")
def _provider_config_response(controller_module: ModuleType) -> tuple[dict, dict]:
expected = {
"type": "secret-input",
@ -308,7 +355,7 @@ def test_builtin_provider_add_passes_payload(
assert response == {"result": "success"}
service_mock.assert_called_once_with(
user_id="user-123",
user_id=user.id,
tenant_id="tenant-456",
provider="openai",
credentials={"api_key": "sk-test"},
@ -397,54 +444,96 @@ def test_builtin_provider_info_uses_core_to_dict_tool_projection(
assert "original_credentials" not in resp
def test_builtin_provider_credentials_get(app: Flask, controller_module, monkeypatch: pytest.MonkeyPatch):
user = _mock_account("user-tenant-cred")
_set_current_account(monkeypatch, controller_module, user, "tenant-cred")
service_payload, expected_response = _credential_response(controller_module)
service_mock = MagicMock(return_value=[service_payload])
monkeypatch.setattr(
controller_module.BuiltinToolManageService,
"get_builtin_tool_provider_credentials",
service_mock,
)
with app.test_request_context("/creds", method="GET"):
resp = controller_module.ToolBuiltinProviderGetCredentialsApi().get(provider="demo")
assert resp == [expected_response]
service_mock.assert_called_once_with(
tenant_id="tenant-cred",
provider_name="demo",
session=ANY,
user=user,
include_credential_ids=None,
)
def test_builtin_provider_credentials_get_reads_repeated_include_ids(
app: Flask, controller_module, monkeypatch: pytest.MonkeyPatch
@pytest.mark.parametrize(
"sqlite_session",
[(Account, Tenant, TenantAccountJoin, BuiltinToolProvider, CredentialPermission)],
indirect=True,
)
def test_builtin_provider_credentials_get(
app: Flask,
controller_module: ModuleType,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
):
user = _mock_account("user-tenant-cred")
credential_payload, expected = _credential_response(controller_module)
service_mock = MagicMock(return_value=[credential_payload])
monkeypatch.setattr(
controller_module.BuiltinToolManageService,
"get_builtin_tool_provider_credentials",
service_mock,
tenant = _persist_workspace(sqlite_session, user, "tenant-cred")
credential = _provider_credential(
tenant_id=tenant.id,
user_id=user.id,
credential_name="Credential",
)
other_tenant = Tenant(name="other-tenant")
other_tenant.id = _stable_uuid("tenant:other-tenant")
other_credential = _provider_credential(
tenant_id=other_tenant.id,
user_id=user.id,
credential_name="Other Tenant Credential",
)
sqlite_session.add_all([credential, other_tenant, other_credential])
sqlite_session.commit()
_set_current_account(monkeypatch, controller_module, user, tenant.id)
with app.test_request_context("/creds?include_credential_ids=cred-1&include_credential_ids=cred-2", method="GET"):
with (
_bind_database_session(sqlite_session),
_mock_credential_encryption(controller_module),
app.test_request_context("/creds", method="GET"),
):
response = controller_module.ToolBuiltinProviderGetCredentialsApi().get(provider="demo")
assert [item["id"] for item in response] == [credential.id]
assert response[0]["name"] == "Credential"
assert response[0]["credentials"] == {"api_key": "[__HIDDEN__]"}
assert response[0]["created_by"] == user.id
assert other_credential.id not in {item["id"] for item in response}
@pytest.mark.parametrize(
"sqlite_session",
[(Account, Tenant, TenantAccountJoin, BuiltinToolProvider, CredentialPermission)],
indirect=True,
)
def test_builtin_provider_credentials_get_reads_repeated_include_ids(
app: Flask,
controller_module: ModuleType,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
):
user = _mock_account("user-tenant-cred")
tenant = _persist_workspace(sqlite_session, user, "tenant-cred")
visible_credential = _provider_credential(
tenant_id=tenant.id,
user_id=user.id,
credential_name="Visible Credential",
)
other_user = _mock_account("other-user")
borrowed_credential = _provider_credential(
tenant_id=tenant.id,
user_id=other_user.id,
credential_name="Borrowed Credential",
visibility=PermissionEnum.ONLY_ME,
)
other_membership = TenantAccountJoin(
tenant_id=tenant.id,
account_id=other_user.id,
role=TenantAccountRole.NORMAL,
)
sqlite_session.add_all([visible_credential, other_user, other_membership, borrowed_credential])
sqlite_session.commit()
request_path = (
f"/creds?include_credential_ids={visible_credential.id}&include_credential_ids={borrowed_credential.id}"
)
with (
_bind_database_session(sqlite_session),
_mock_credential_encryption(controller_module),
app.test_request_context(request_path, method="GET"),
):
api = controller_module.ToolBuiltinProviderGetCredentialsApi()
resp = unwrap(api.get)(api, "tenant-cred", user, provider="demo")
response = unwrap(api.get)(api, tenant.id, user, provider="demo")
assert resp == [expected]
service_mock.assert_called_once_with(
tenant_id="tenant-cred",
provider_name="demo",
session=ANY,
user=user,
include_credential_ids=["cred-1", "cred-2"],
)
assert [item["id"] for item in response] == [visible_credential.id, borrowed_credential.id]
assert response[0]["from_other_member"] is False
assert response[1]["from_other_member"] is True
def test_api_provider_remote_schema_get(app: Flask, controller_module, monkeypatch: pytest.MonkeyPatch):