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 import importlib
from contextlib import ExitStack, contextmanager from contextlib import ExitStack, contextmanager
from inspect import unwrap from inspect import unwrap
from types import ModuleType, SimpleNamespace from types import ModuleType
from unittest.mock import ANY, MagicMock, patch from unittest.mock import MagicMock, patch
from uuid import NAMESPACE_URL, uuid5
import pytest import pytest
from flask import Flask from flask import Flask
from flask.views import MethodView 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.api_entities import ToolProviderApiEntity as CoreToolProviderApiEntity
from core.tools.entities.common_entities import I18nObject from core.tools.entities.common_entities import I18nObject
from core.tools.entities.tool_entities import ToolParameter 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.account import TenantAccountRole
from models.credential_permission import CredentialPermission
from models.enums import PermissionEnum
if not hasattr(builtins, "MethodView"): if not hasattr(builtins, "MethodView"):
builtins.MethodView = MethodView # type: ignore[attr-defined] builtins.MethodView = MethodView # type: ignore[attr-defined]
@ -27,13 +31,6 @@ _CONTROLLER_MODULE: ModuleType | None = None
_WRAPS_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 @pytest.fixture
def app() -> Flask: def app() -> Flask:
flask_app = Flask(__name__) flask_app = Flask(__name__)
@ -69,8 +66,7 @@ def controller_module(monkeypatch: pytest.MonkeyPatch):
with ExitStack() as stack: with ExitStack() as stack:
for target, value in patch_targets: for target, value in patch_targets:
stack.enter_context(patch(target, value)) 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 module = _CONTROLLER_MODULE
@ -88,11 +84,79 @@ def controller_module(monkeypatch: pytest.MonkeyPatch):
def _mock_account(user_id: str = "user-123") -> Account: def _mock_account(user_id: str = "user-123") -> Account:
user = Account(name="Test User", email=f"{user_id}@example.com") 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 user.role = TenantAccountRole.NORMAL
return user 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( def _set_current_account(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
controller_module: ModuleType, controller_module: ModuleType,
@ -190,23 +254,6 @@ def _provider_list_item(
return service_payload, provider.model_dump(mode="json", exclude_unset=True) 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]: def _provider_config_response(controller_module: ModuleType) -> tuple[dict, dict]:
expected = { expected = {
"type": "secret-input", "type": "secret-input",
@ -308,7 +355,7 @@ def test_builtin_provider_add_passes_payload(
assert response == {"result": "success"} assert response == {"result": "success"}
service_mock.assert_called_once_with( service_mock.assert_called_once_with(
user_id="user-123", user_id=user.id,
tenant_id="tenant-456", tenant_id="tenant-456",
provider="openai", provider="openai",
credentials={"api_key": "sk-test"}, 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 assert "original_credentials" not in resp
def test_builtin_provider_credentials_get(app: Flask, controller_module, monkeypatch: pytest.MonkeyPatch): @pytest.mark.parametrize(
user = _mock_account("user-tenant-cred") "sqlite_session",
_set_current_account(monkeypatch, controller_module, user, "tenant-cred") [(Account, Tenant, TenantAccountJoin, BuiltinToolProvider, CredentialPermission)],
service_payload, expected_response = _credential_response(controller_module) indirect=True,
service_mock = MagicMock(return_value=[service_payload]) )
monkeypatch.setattr( def test_builtin_provider_credentials_get(
controller_module.BuiltinToolManageService, app: Flask,
"get_builtin_tool_provider_credentials", controller_module: ModuleType,
service_mock, monkeypatch: pytest.MonkeyPatch,
) sqlite_session: Session,
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
): ):
user = _mock_account("user-tenant-cred") user = _mock_account("user-tenant-cred")
credential_payload, expected = _credential_response(controller_module) tenant = _persist_workspace(sqlite_session, user, "tenant-cred")
service_mock = MagicMock(return_value=[credential_payload]) credential = _provider_credential(
monkeypatch.setattr( tenant_id=tenant.id,
controller_module.BuiltinToolManageService, user_id=user.id,
"get_builtin_tool_provider_credentials", credential_name="Credential",
service_mock,
) )
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() 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] assert [item["id"] for item in response] == [visible_credential.id, borrowed_credential.id]
service_mock.assert_called_once_with( assert response[0]["from_other_member"] is False
tenant_id="tenant-cred", assert response[1]["from_other_member"] is True
provider_name="demo",
session=ANY,
user=user,
include_credential_ids=["cred-1", "cred-2"],
)
def test_api_provider_remote_schema_get(app: Flask, controller_module, monkeypatch: pytest.MonkeyPatch): def test_api_provider_remote_schema_get(app: Flask, controller_module, monkeypatch: pytest.MonkeyPatch):