mirror of
https://github.com/langgenius/dify.git
synced 2026-07-26 06:08:38 +08:00
test: use sqlite3 session in test_datasource_provider_service (#38695)
This commit is contained in:
parent
a2d9aeff37
commit
8906a49e56
@ -1,15 +1,20 @@
|
||||
from collections.abc import Iterator
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy import Engine, select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.plugin.entities.plugin_daemon import CredentialType
|
||||
from graphon.model_runtime.entities.provider_entities import FormType
|
||||
from models.account import Account
|
||||
from models.base import TypeBase
|
||||
from models.model import EndUser
|
||||
from models.oauth import DatasourceProvider
|
||||
from models.oauth import DatasourceOauthParamConfig, DatasourceOauthTenantParamConfig, DatasourceProvider
|
||||
from models.provider_ids import DatasourceProviderID
|
||||
from services import datasource_provider_service as service_module
|
||||
from services.datasource_provider_service import DatasourceProviderService, get_current_user
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@ -21,6 +26,37 @@ def make_id(s: str = "org/plugin/provider") -> DatasourceProviderID:
|
||||
return DatasourceProviderID(s)
|
||||
|
||||
|
||||
def make_provider(
|
||||
*,
|
||||
credential_id: str = "cred-id",
|
||||
tenant_id: str = "t1",
|
||||
name: str = "name",
|
||||
provider: str = "prov",
|
||||
plugin_id: str = "org/plug",
|
||||
auth_type: str = "api_key",
|
||||
encrypted_credentials: dict[str, object] | None = None,
|
||||
is_default: bool = False,
|
||||
expires_at: int = -1,
|
||||
) -> DatasourceProvider:
|
||||
datasource_provider = DatasourceProvider(
|
||||
tenant_id=tenant_id,
|
||||
name=name,
|
||||
provider=provider,
|
||||
plugin_id=plugin_id,
|
||||
auth_type=auth_type,
|
||||
encrypted_credentials=encrypted_credentials or {},
|
||||
is_default=is_default,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
datasource_provider.id = credential_id
|
||||
return datasource_provider
|
||||
|
||||
|
||||
def persist(session: Session, *models: TypeBase) -> None:
|
||||
session.add_all(models)
|
||||
session.commit()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test class
|
||||
# ---------------------------------------------------------------------------
|
||||
@ -34,33 +70,17 @@ class TestDatasourceProviderService:
|
||||
return DatasourceProviderService()
|
||||
|
||||
@pytest.fixture
|
||||
def mock_db_session(self):
|
||||
"""
|
||||
Mock session with scalar/scalars defaults for current SQLAlchemy access paths.
|
||||
"""
|
||||
with (
|
||||
patch("services.datasource_provider_service.Session") as mock_cls,
|
||||
patch("services.datasource_provider_service.sessionmaker") as mock_sm,
|
||||
):
|
||||
sess = MagicMock(spec=Session)
|
||||
|
||||
# Default values for select()-style calls (tests override per-case)
|
||||
sess.scalar.return_value = None
|
||||
sess.scalars.return_value.all.return_value = []
|
||||
|
||||
mock_cls.return_value.__enter__.return_value = sess
|
||||
mock_cls.return_value.no_autoflush.__enter__.return_value = sess
|
||||
mock_sm.return_value.begin.return_value.__enter__.return_value = sess
|
||||
mock_sm.return_value.begin.return_value.__exit__ = MagicMock(return_value=False)
|
||||
|
||||
yield sess
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_db(self, mock_db_session):
|
||||
with patch("services.datasource_provider_service.db") as mock_db:
|
||||
mock_db.session = mock_db_session
|
||||
mock_db.engine = MagicMock()
|
||||
yield mock_db
|
||||
def sqlite_session(self, sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]:
|
||||
"""Provide the service with an isolated SQLite-backed session and engine."""
|
||||
tables = [
|
||||
TypeBase.metadata.tables[model.__tablename__]
|
||||
for model in (DatasourceOauthParamConfig, DatasourceOauthTenantParamConfig, DatasourceProvider)
|
||||
]
|
||||
TypeBase.metadata.create_all(sqlite_engine, tables=tables)
|
||||
session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
|
||||
with session_factory() as session:
|
||||
monkeypatch.setattr(service_module, "db", SimpleNamespace(engine=sqlite_engine, session=session))
|
||||
yield session
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_externals(self):
|
||||
@ -162,12 +182,18 @@ class TestDatasourceProviderService:
|
||||
# is_system_oauth_params_exist (line 357-363)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_true_when_system_oauth_params_exist(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = MagicMock()
|
||||
def test_should_return_true_when_system_oauth_params_exist(self, service, sqlite_session):
|
||||
persist(
|
||||
sqlite_session,
|
||||
DatasourceOauthParamConfig(
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
system_credentials={"client_id": "configured"},
|
||||
),
|
||||
)
|
||||
assert service.is_system_oauth_params_exist(make_id()) is True
|
||||
|
||||
def test_should_return_false_when_system_oauth_params_missing(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = None
|
||||
def test_should_return_false_when_system_oauth_params_missing(self, service, sqlite_session):
|
||||
assert service.is_system_oauth_params_exist(make_id()) is False
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
@ -175,59 +201,93 @@ class TestDatasourceProviderService:
|
||||
# NOTE: uses .count() not .first()
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_true_when_tenant_oauth_params_enabled(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = 1
|
||||
assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=mock_db_session) is True
|
||||
def test_should_return_true_when_tenant_oauth_params_enabled(self, service, sqlite_session):
|
||||
persist(
|
||||
sqlite_session,
|
||||
DatasourceOauthTenantParamConfig(
|
||||
tenant_id="t1",
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
enabled=True,
|
||||
),
|
||||
)
|
||||
assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=sqlite_session) is True
|
||||
|
||||
def test_should_return_false_when_tenant_oauth_params_disabled(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = 0
|
||||
assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=mock_db_session) is False
|
||||
def test_should_return_false_when_tenant_oauth_params_disabled(self, service, sqlite_session):
|
||||
persist(
|
||||
sqlite_session,
|
||||
DatasourceOauthTenantParamConfig(
|
||||
tenant_id="t1",
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
enabled=False,
|
||||
),
|
||||
)
|
||||
assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=sqlite_session) is False
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# remove_oauth_custom_client_params (lines 55-61)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_delete_tenant_config_when_removing_oauth_params(self, service, mock_db_session):
|
||||
def test_should_delete_tenant_config_when_removing_oauth_params(self, service, sqlite_session):
|
||||
config = DatasourceOauthTenantParamConfig(
|
||||
tenant_id="t1",
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
enabled=True,
|
||||
)
|
||||
persist(sqlite_session, config)
|
||||
config_id = config.id
|
||||
service.remove_oauth_custom_client_params("t1", make_id())
|
||||
mock_db_session.execute.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(DatasourceOauthTenantParamConfig, config_id) is None
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# setup_oauth_custom_client_params (315-351)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_skip_db_write_when_credentials_are_none(self, service, mock_db_session):
|
||||
def test_should_skip_db_write_when_credentials_are_none(self, service, sqlite_session):
|
||||
"""When credentials=None, should return immediately without any DB write."""
|
||||
service.setup_oauth_custom_client_params("t1", make_id(), None, None)
|
||||
mock_db_session.add.assert_not_called()
|
||||
assert sqlite_session.scalars(select(DatasourceOauthTenantParamConfig)).all() == []
|
||||
|
||||
def test_should_create_new_config_when_none_exists(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = None
|
||||
def test_should_create_new_config_when_none_exists(self, service, sqlite_session):
|
||||
with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)):
|
||||
service.setup_oauth_custom_client_params("t1", make_id(), {"k": "v"}, True)
|
||||
mock_db_session.add.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
config = sqlite_session.scalar(select(DatasourceOauthTenantParamConfig))
|
||||
assert config is not None
|
||||
assert config.tenant_id == "t1"
|
||||
assert config.enabled is True
|
||||
assert config.client_params == {"k": "enc"}
|
||||
|
||||
def test_should_update_existing_config_when_record_found(self, service, mock_db_session):
|
||||
existing = MagicMock()
|
||||
mock_db_session.scalar.return_value = existing
|
||||
def test_should_update_existing_config_when_record_found(self, service, sqlite_session):
|
||||
existing = DatasourceOauthTenantParamConfig(
|
||||
tenant_id="t1",
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
client_params={"k": "old"},
|
||||
enabled=True,
|
||||
)
|
||||
persist(sqlite_session, existing)
|
||||
with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)):
|
||||
service.setup_oauth_custom_client_params("t1", make_id(), {"k": "v"}, False)
|
||||
mock_db_session.add.assert_not_called() # update in place, no add
|
||||
sqlite_session.refresh(existing)
|
||||
assert existing.client_params == {"k": "enc"}
|
||||
assert existing.enabled is False
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# decrypt / encrypt credentials (lines 70-98)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_decrypt_secret_fields_when_decrypting_api_key_credentials(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "api_key"
|
||||
p.encrypted_credentials = {"sk": "enc_val"}
|
||||
def test_should_decrypt_secret_fields_when_decrypting_api_key_credentials(self, service, sqlite_session):
|
||||
p = make_provider(encrypted_credentials={"sk": "enc_val"})
|
||||
with patch.object(service, "extract_secret_variables", return_value=["sk"]):
|
||||
result = service.decrypt_datasource_provider_credentials("t1", p, "org/plug", "prov")
|
||||
assert result["sk"] == "dec_tok"
|
||||
|
||||
def test_should_encrypt_secret_fields_when_encrypting_api_key_credentials(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "api_key"
|
||||
def test_should_encrypt_secret_fields_when_encrypting_api_key_credentials(self, service, sqlite_session):
|
||||
p = make_provider()
|
||||
with patch.object(service, "extract_secret_variables", return_value=["sk"]):
|
||||
result = service.encrypt_datasource_provider_credentials("t1", "prov", "org/plug", {"sk": "plain"}, p)
|
||||
assert result["sk"] == "enc_tok"
|
||||
@ -237,33 +297,31 @@ class TestDatasourceProviderService:
|
||||
# get_datasource_credentials (lines 113-165)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_empty_dict_when_credential_not_found(self, service, mock_db_session, mock_user):
|
||||
def test_should_return_empty_dict_when_credential_not_found(self, service, sqlite_session, mock_user):
|
||||
with patch("services.datasource_provider_service.get_current_user", return_value=mock_user):
|
||||
mock_db_session.scalar.return_value = None
|
||||
assert service.get_datasource_credentials("t1", "prov", "org/plug") == {}
|
||||
|
||||
def test_should_refresh_oauth_tokens_when_expired(self, service, mock_db_session, mock_user):
|
||||
def test_should_refresh_oauth_tokens_when_expired(self, service, sqlite_session, mock_user):
|
||||
"""Expired OAuth credential (expires_at near zero) triggers a refresh."""
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "oauth2"
|
||||
p.expires_at = 0 # expired
|
||||
p.encrypted_credentials = {"tok": "x"}
|
||||
mock_db_session.scalar.return_value = p
|
||||
p = make_provider(auth_type="oauth2", expires_at=0, encrypted_credentials={"tok": "x"})
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service, "get_oauth_client", return_value={"oc": "v"}),
|
||||
patch.object(service, "decrypt_datasource_provider_credentials", return_value={"tok": "plain"}),
|
||||
):
|
||||
service.get_datasource_credentials("t1", "prov", "org/plug")
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(DatasourceProvider, p.id).expires_at == 9999
|
||||
|
||||
def test_should_include_provider_name_when_refresh_fails(self, service, mock_db_session, mock_user):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.id = "cred-id"
|
||||
p.name = "Credential"
|
||||
p.auth_type = "oauth2"
|
||||
p.expires_at = 0
|
||||
p.encrypted_credentials = {"tok": "x"}
|
||||
mock_db_session.scalar.return_value = p
|
||||
def test_should_include_provider_name_when_refresh_fails(self, service, sqlite_session, mock_user):
|
||||
p = make_provider(
|
||||
name="Credential",
|
||||
auth_type="oauth2",
|
||||
expires_at=0,
|
||||
encrypted_credentials={"tok": "x"},
|
||||
)
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch("services.datasource_provider_service.OAuthHandler") as oauth_handler,
|
||||
@ -274,13 +332,10 @@ class TestDatasourceProviderService:
|
||||
with pytest.raises(ValueError, match="provider prov"):
|
||||
service.get_datasource_credentials("t1", "prov", "org/plug")
|
||||
|
||||
def test_should_return_decrypted_credentials_when_api_key_not_expired(self, service, mock_db_session, mock_user):
|
||||
def test_should_return_decrypted_credentials_when_api_key_not_expired(self, service, sqlite_session, mock_user):
|
||||
"""API key credentials with expires_at=-1 skip refresh and return directly."""
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "api_key"
|
||||
p.expires_at = -1 # sentinel: never expires
|
||||
p.encrypted_credentials = {"k": "v"}
|
||||
mock_db_session.scalar.return_value = p
|
||||
p = make_provider(encrypted_credentials={"k": "v"})
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service, "decrypt_datasource_provider_credentials", return_value={"k": "plain"}),
|
||||
@ -288,13 +343,10 @@ class TestDatasourceProviderService:
|
||||
result = service.get_datasource_credentials("t1", "prov", "org/plug")
|
||||
assert result == {"k": "plain"}
|
||||
|
||||
def test_should_fetch_by_credential_id_when_provided(self, service, mock_db_session, mock_user):
|
||||
def test_should_fetch_by_credential_id_when_provided(self, service, sqlite_session, mock_user):
|
||||
"""When credential_id is passed, the credential_id filter path (line 113) is taken."""
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "api_key"
|
||||
p.expires_at = -1
|
||||
p.encrypted_credentials = {}
|
||||
mock_db_session.scalar.return_value = p
|
||||
p = make_provider(credential_id="cred-id", provider="other-provider", plugin_id="other-plugin")
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service, "decrypt_datasource_provider_credentials", return_value={"k": "v"}),
|
||||
@ -306,17 +358,13 @@ class TestDatasourceProviderService:
|
||||
# get_all_datasource_credentials_by_provider (lines 176-228)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_empty_list_when_no_provider_credentials_exist(self, service, mock_db_session, mock_user):
|
||||
def test_should_return_empty_list_when_no_provider_credentials_exist(self, service, sqlite_session, mock_user):
|
||||
with patch("services.datasource_provider_service.get_current_user", return_value=mock_user):
|
||||
mock_db_session.scalars.return_value.all.return_value = []
|
||||
assert service.get_all_datasource_credentials_by_provider("t1", "prov", "org/plug") == []
|
||||
|
||||
def test_should_refresh_and_return_credentials_when_oauth_expired(self, service, mock_db_session, mock_user):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "oauth2"
|
||||
p.expires_at = 0
|
||||
p.encrypted_credentials = {"t": "x"}
|
||||
mock_db_session.scalars.return_value.all.return_value = [p]
|
||||
def test_should_refresh_and_return_credentials_when_oauth_expired(self, service, sqlite_session, mock_user):
|
||||
p = make_provider(auth_type="oauth2", expires_at=0, encrypted_credentials={"t": "x"})
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service, "get_oauth_client", return_value={"oc": "v"}),
|
||||
@ -326,19 +374,21 @@ class TestDatasourceProviderService:
|
||||
assert len(result) == 1
|
||||
|
||||
def test_should_skip_failed_provider_when_refreshing_all_credentials(
|
||||
self, service, mock_db_session, mock_user, caplog
|
||||
self, service, sqlite_session, mock_user, caplog
|
||||
):
|
||||
failed_provider = MagicMock(spec=DatasourceProvider)
|
||||
failed_provider.id = "failed-cred"
|
||||
failed_provider.name = "Failed"
|
||||
failed_provider.auth_type = "oauth2"
|
||||
failed_provider.expires_at = 0
|
||||
working_provider = MagicMock(spec=DatasourceProvider)
|
||||
working_provider.id = "working-cred"
|
||||
working_provider.name = "Working"
|
||||
working_provider.auth_type = "oauth2"
|
||||
working_provider.expires_at = 0
|
||||
mock_db_session.scalars.return_value.all.return_value = [failed_provider, working_provider]
|
||||
failed_provider = make_provider(
|
||||
credential_id="failed-cred",
|
||||
name="Failed",
|
||||
auth_type="oauth2",
|
||||
expires_at=0,
|
||||
)
|
||||
working_provider = make_provider(
|
||||
credential_id="working-cred",
|
||||
name="Working",
|
||||
auth_type="oauth2",
|
||||
expires_at=0,
|
||||
)
|
||||
persist(sqlite_session, failed_provider, working_provider)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(
|
||||
@ -354,13 +404,10 @@ class TestDatasourceProviderService:
|
||||
assert "Skipping datasource credentials for provider prov" in caplog.text
|
||||
|
||||
def test_should_return_valid_credentials_without_refresh_when_getting_all_credentials(
|
||||
self, service, mock_db_session, mock_user
|
||||
self, service, sqlite_session, mock_user
|
||||
):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "oauth2"
|
||||
p.expires_at = -1
|
||||
p.encrypted_credentials = {"t": "x"}
|
||||
mock_db_session.scalars.return_value.all.return_value = [p]
|
||||
p = make_provider(auth_type="oauth2", encrypted_credentials={"t": "x"})
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service, "_refresh_datasource_credentials") as refresh_credentials,
|
||||
@ -374,49 +421,78 @@ class TestDatasourceProviderService:
|
||||
# update_datasource_provider_name (lines 236-303)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_raise_value_error_when_provider_not_found_on_name_update(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = None
|
||||
def test_should_raise_value_error_when_provider_not_found_on_name_update(self, service, sqlite_session):
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
service.update_datasource_provider_name("t1", make_id(), "new", "cred-id")
|
||||
|
||||
def test_should_return_early_when_new_name_matches_current(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.name = "same"
|
||||
mock_db_session.scalar.return_value = p
|
||||
def test_should_return_early_when_new_name_matches_current(self, service, sqlite_session):
|
||||
p = make_provider(
|
||||
credential_id="cred-id",
|
||||
name="same",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
)
|
||||
persist(sqlite_session, p)
|
||||
service.update_datasource_provider_name("t1", make_id(), "same", "cred-id")
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(DatasourceProvider, p.id).name == "same"
|
||||
|
||||
def test_should_raise_value_error_when_name_already_exists(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.name = "old_name"
|
||||
p.is_default = False
|
||||
mock_db_session.scalar.side_effect = [p, 1] # first: fetch provider, second: name conflict count
|
||||
def test_should_raise_value_error_when_name_already_exists(self, service, sqlite_session):
|
||||
p = make_provider(
|
||||
credential_id="some-id",
|
||||
name="old_name",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
)
|
||||
conflict = make_provider(
|
||||
credential_id="conflict-id",
|
||||
name="new_name",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
)
|
||||
persist(sqlite_session, p, conflict)
|
||||
with pytest.raises(ValueError, match="already exists"):
|
||||
service.update_datasource_provider_name("t1", make_id(), "new_name", "some-id")
|
||||
|
||||
def test_should_update_name_and_commit_when_no_conflict(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.name = "old_name"
|
||||
p.is_default = False
|
||||
mock_db_session.scalar.side_effect = [p, 0] # first: fetch provider, second: name conflict count
|
||||
def test_should_update_name_and_commit_when_no_conflict(self, service, sqlite_session):
|
||||
p = make_provider(
|
||||
credential_id="some-id",
|
||||
name="old_name",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
)
|
||||
persist(sqlite_session, p)
|
||||
service.update_datasource_provider_name("t1", make_id(), "new_name", "some-id")
|
||||
assert p.name == "new_name"
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(DatasourceProvider, p.id).name == "new_name"
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# set_default_datasource_provider (lines 277-303)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_raise_value_error_when_target_provider_not_found(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = None
|
||||
def test_should_raise_value_error_when_target_provider_not_found(self, service, sqlite_session):
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
service.set_default_datasource_provider("t1", make_id(), "bad-id")
|
||||
|
||||
def test_should_mark_target_as_default_and_commit(self, service, mock_db_session):
|
||||
target = MagicMock(spec=DatasourceProvider)
|
||||
target.provider = "provider"
|
||||
target.plugin_id = "org/plug"
|
||||
mock_db_session.scalar.return_value = target
|
||||
def test_should_mark_target_as_default_and_commit(self, service, sqlite_session):
|
||||
current_default = make_provider(
|
||||
credential_id="old-id",
|
||||
name="old",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
is_default=True,
|
||||
)
|
||||
target = make_provider(
|
||||
credential_id="new-id",
|
||||
name="new",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
)
|
||||
persist(sqlite_session, current_default, target)
|
||||
service.set_default_datasource_provider("t1", make_id(), "new-id")
|
||||
assert target.is_default is True
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(DatasourceProvider, current_default.id).is_default is False
|
||||
assert sqlite_session.get(DatasourceProvider, target.id).is_default is True
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# get_oauth_encrypter (lines 404-420)
|
||||
@ -448,38 +524,61 @@ class TestDatasourceProviderService:
|
||||
# get_tenant_oauth_client (lines 381-402)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_masked_credentials_when_mask_is_true(self, service, mock_db_session):
|
||||
tenant_params = MagicMock()
|
||||
tenant_params.client_params = {"k": "v"}
|
||||
mock_db_session.scalar.return_value = tenant_params
|
||||
def test_should_return_masked_credentials_when_mask_is_true(self, service, sqlite_session):
|
||||
tenant_params = DatasourceOauthTenantParamConfig(
|
||||
tenant_id="t1",
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
client_params={"k": "v"},
|
||||
)
|
||||
persist(sqlite_session, tenant_params)
|
||||
with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)):
|
||||
result = service.get_tenant_oauth_client("t1", make_id(), mask=True, session=mock_db_session)
|
||||
result = service.get_tenant_oauth_client("t1", make_id(), mask=True, session=sqlite_session)
|
||||
assert result == {"k": "mask"}
|
||||
|
||||
def test_should_return_decrypted_credentials_when_mask_is_false(self, service, mock_db_session):
|
||||
tenant_params = MagicMock()
|
||||
tenant_params.client_params = {"k": "v"}
|
||||
mock_db_session.scalar.return_value = tenant_params
|
||||
def test_should_return_decrypted_credentials_when_mask_is_false(self, service, sqlite_session):
|
||||
tenant_params = DatasourceOauthTenantParamConfig(
|
||||
tenant_id="t1",
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
client_params={"k": "v"},
|
||||
)
|
||||
persist(sqlite_session, tenant_params)
|
||||
with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)):
|
||||
result = service.get_tenant_oauth_client("t1", make_id(), mask=False, session=mock_db_session)
|
||||
result = service.get_tenant_oauth_client("t1", make_id(), mask=False, session=sqlite_session)
|
||||
assert result == {"k": "dec"}
|
||||
|
||||
def test_should_return_none_when_no_tenant_oauth_config_exists(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = None
|
||||
assert service.get_tenant_oauth_client("t1", make_id(), session=mock_db_session) is None
|
||||
def test_should_return_none_when_no_tenant_oauth_config_exists(self, service, sqlite_session):
|
||||
assert service.get_tenant_oauth_client("t1", make_id(), session=sqlite_session) is None
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# get_oauth_client (lines 423-457)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_use_tenant_config_when_available(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = MagicMock(client_params={"k": "v"})
|
||||
def test_should_use_tenant_config_when_available(self, service, sqlite_session):
|
||||
persist(
|
||||
sqlite_session,
|
||||
DatasourceOauthTenantParamConfig(
|
||||
tenant_id="t1",
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
client_params={"k": "v"},
|
||||
enabled=True,
|
||||
),
|
||||
)
|
||||
with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)):
|
||||
result = service.get_oauth_client("t1", make_id())
|
||||
assert result == {"k": "dec"}
|
||||
|
||||
def test_should_fallback_to_system_credentials_when_tenant_config_missing(self, service, mock_db_session):
|
||||
mock_db_session.scalar.side_effect = [None, MagicMock(system_credentials={"k": "sys"})]
|
||||
def test_should_fallback_to_system_credentials_when_tenant_config_missing(self, service, sqlite_session):
|
||||
persist(
|
||||
sqlite_session,
|
||||
DatasourceOauthParamConfig(
|
||||
plugin_id="org/plugin",
|
||||
provider="provider",
|
||||
system_credentials={"k": "sys"},
|
||||
),
|
||||
)
|
||||
with (
|
||||
patch.object(service.provider_manager, "fetch_datasource_provider"),
|
||||
patch("services.datasource_provider_service.PluginService.is_plugin_verified", return_value=True),
|
||||
@ -487,9 +586,8 @@ class TestDatasourceProviderService:
|
||||
result = service.get_oauth_client("t1", make_id())
|
||||
assert result == {"k": "sys"}
|
||||
|
||||
def test_should_raise_value_error_when_no_oauth_config_available(self, service, mock_db_session):
|
||||
def test_should_raise_value_error_when_no_oauth_config_available(self, service, sqlite_session):
|
||||
"""Neither tenant nor system credentials → raises ValueError."""
|
||||
mock_db_session.scalar.side_effect = [None, None]
|
||||
with (
|
||||
patch.object(service.provider_manager, "fetch_datasource_provider"),
|
||||
patch("services.datasource_provider_service.PluginService.is_plugin_verified", return_value=False),
|
||||
@ -501,40 +599,53 @@ class TestDatasourceProviderService:
|
||||
# add_datasource_oauth_provider (lines 539-607)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_add_oauth_provider_successfully_when_name_is_unique(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = 0
|
||||
def test_should_add_oauth_provider_successfully_when_name_is_unique(self, service, sqlite_session):
|
||||
with patch.object(service, "extract_secret_variables", return_value=[]):
|
||||
service.add_datasource_oauth_provider("new", "t1", make_id(), "http://cb", 9999, {})
|
||||
mock_db_session.add.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
provider = sqlite_session.scalar(select(DatasourceProvider))
|
||||
assert provider is not None
|
||||
assert provider.name == "new"
|
||||
assert provider.auth_type == CredentialType.OAUTH2.value
|
||||
|
||||
def test_should_auto_rename_when_oauth_provider_name_conflicts(self, service, mock_db_session):
|
||||
def test_should_auto_rename_when_oauth_provider_name_conflicts(self, service, sqlite_session):
|
||||
"""Conflict on name results in auto-incremented name, not an error."""
|
||||
mock_db_session.scalar.return_value = 1 # conflict first, then auto-named
|
||||
persist(
|
||||
sqlite_session,
|
||||
make_provider(
|
||||
name="conflict",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
auth_type=CredentialType.OAUTH2.value,
|
||||
),
|
||||
)
|
||||
with (
|
||||
patch.object(service, "extract_secret_variables", return_value=[]),
|
||||
patch.object(service, "generate_next_datasource_provider_name", return_value="new_gen"),
|
||||
):
|
||||
service.add_datasource_oauth_provider("conflict", "t1", make_id(), "http://cb", 9999, {})
|
||||
mock_db_session.add.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
names = set(sqlite_session.scalars(select(DatasourceProvider.name)).all())
|
||||
assert names == {"conflict", "gen_name"}
|
||||
|
||||
def test_should_auto_generate_name_when_none_provided_for_oauth(self, service, mock_db_session):
|
||||
def test_should_auto_generate_name_when_none_provided_for_oauth(self, service, sqlite_session):
|
||||
"""name=None causes auto-generation via generate_next_datasource_provider_name."""
|
||||
mock_db_session.scalar.return_value = 0
|
||||
with (
|
||||
patch.object(service, "extract_secret_variables", return_value=[]),
|
||||
patch.object(service, "generate_next_datasource_provider_name", return_value="auto"),
|
||||
):
|
||||
service.add_datasource_oauth_provider(None, "t1", make_id(), "http://cb", 9999, {})
|
||||
mock_db_session.add.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.scalar(select(DatasourceProvider.name)) == "auto"
|
||||
|
||||
def test_should_encrypt_secret_fields_when_adding_oauth_provider(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = 0
|
||||
def test_should_encrypt_secret_fields_when_adding_oauth_provider(self, service, sqlite_session):
|
||||
with patch.object(service, "extract_secret_variables", return_value=["secret_key"]):
|
||||
service.add_datasource_oauth_provider("nm", "t1", make_id(), "http://cb", 9999, {"secret_key": "value"})
|
||||
self._enc.encrypt_token.assert_called()
|
||||
sqlite_session.expire_all()
|
||||
provider = sqlite_session.scalar(select(DatasourceProvider))
|
||||
assert provider.encrypted_credentials == {"secret_key": "enc_tok"}
|
||||
|
||||
def test_should_acquire_redis_lock_when_adding_oauth_provider(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = 0
|
||||
def test_should_acquire_redis_lock_when_adding_oauth_provider(self, service, sqlite_session):
|
||||
with patch.object(service, "extract_secret_variables", return_value=[]):
|
||||
service.add_datasource_oauth_provider("nm", "t1", make_id(), "http://cb", 9999, {})
|
||||
self._redis.lock.assert_called()
|
||||
@ -543,37 +654,71 @@ class TestDatasourceProviderService:
|
||||
# reauthorize_datasource_oauth_provider (lines 477-537)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_raise_value_error_when_credential_id_not_found_on_reauth(self, service, mock_db_session):
|
||||
mock_db_session.scalar.return_value = None
|
||||
def test_should_raise_value_error_when_credential_id_not_found_on_reauth(self, service, sqlite_session):
|
||||
with patch.object(service, "extract_secret_variables", return_value=[]):
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
service.reauthorize_datasource_oauth_provider("n", "t1", make_id(), "u", 1, {}, "bad-id")
|
||||
|
||||
def test_should_reauthorize_and_commit_when_credential_found(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
mock_db_session.scalar.side_effect = [p, 0] # first: fetch provider, second: name conflict count
|
||||
def test_should_reauthorize_and_commit_when_credential_found(self, service, sqlite_session):
|
||||
p = make_provider(
|
||||
credential_id="oid",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
auth_type=CredentialType.OAUTH2.value,
|
||||
)
|
||||
persist(sqlite_session, p)
|
||||
with patch.object(service, "extract_secret_variables", return_value=[]):
|
||||
service.reauthorize_datasource_oauth_provider("n", "t1", make_id(), "u", 1, {}, "oid")
|
||||
sqlite_session.expire_all()
|
||||
updated = sqlite_session.get(DatasourceProvider, p.id)
|
||||
assert updated.expires_at == 1
|
||||
assert updated.avatar_url == "u"
|
||||
|
||||
def test_should_auto_rename_when_reauth_name_conflicts(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
mock_db_session.scalar.side_effect = [p, 1] # first: fetch provider, second: name conflict count
|
||||
mock_db_session.scalars.return_value.all.return_value = []
|
||||
def test_should_auto_rename_when_reauth_name_conflicts(self, service, sqlite_session):
|
||||
p = make_provider(
|
||||
credential_id="cred-id",
|
||||
name="original",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
auth_type=CredentialType.OAUTH2.value,
|
||||
)
|
||||
conflict = make_provider(
|
||||
credential_id="conflict-id",
|
||||
name="conflict_name",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
auth_type=CredentialType.OAUTH2.value,
|
||||
)
|
||||
persist(sqlite_session, p, conflict)
|
||||
with patch.object(service, "extract_secret_variables", return_value=["tok"]):
|
||||
service.reauthorize_datasource_oauth_provider(
|
||||
"conflict_name", "t1", make_id(), "u", 9999, {"tok": "v"}, "cred-id"
|
||||
)
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(DatasourceProvider, p.id).encrypted_credentials == {"tok": "enc_tok"}
|
||||
|
||||
def test_should_encrypt_secret_fields_when_reauthorizing(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
mock_db_session.scalar.side_effect = [p, 0] # first: fetch provider, second: name conflict count
|
||||
def test_should_encrypt_secret_fields_when_reauthorizing(self, service, sqlite_session):
|
||||
p = make_provider(
|
||||
credential_id="cred-id",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
auth_type=CredentialType.OAUTH2.value,
|
||||
)
|
||||
persist(sqlite_session, p)
|
||||
with patch.object(service, "extract_secret_variables", return_value=["tok"]):
|
||||
service.reauthorize_datasource_oauth_provider(None, "t1", make_id(), "u", 9999, {"tok": "val"}, "cred-id")
|
||||
self._enc.encrypt_token.assert_called()
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(DatasourceProvider, p.id).encrypted_credentials == {"tok": "enc_tok"}
|
||||
|
||||
def test_should_acquire_redis_lock_when_reauthorizing(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
mock_db_session.scalar.side_effect = [p, 0] # first: fetch provider, second: name conflict count
|
||||
def test_should_acquire_redis_lock_when_reauthorizing(self, service, sqlite_session):
|
||||
p = make_provider(
|
||||
credential_id="oid",
|
||||
provider="provider",
|
||||
plugin_id="org/plugin",
|
||||
auth_type=CredentialType.OAUTH2.value,
|
||||
)
|
||||
persist(sqlite_session, p)
|
||||
with patch.object(service, "extract_secret_variables", return_value=[]):
|
||||
service.reauthorize_datasource_oauth_provider("n", "t1", make_id(), "u", 1, {}, "oid")
|
||||
self._redis.lock.assert_called()
|
||||
@ -582,15 +727,17 @@ class TestDatasourceProviderService:
|
||||
# add_datasource_api_key_provider (lines 608-675)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_raise_value_error_when_api_key_name_already_exists(self, service, mock_db_session, mock_user):
|
||||
def test_should_raise_value_error_when_api_key_name_already_exists(self, service, sqlite_session, mock_user):
|
||||
"""explicit name supplied + conflict → raises ValueError immediately."""
|
||||
mock_db_session.scalar.return_value = 1
|
||||
persist(
|
||||
sqlite_session,
|
||||
make_provider(name="clash", provider="provider", plugin_id="org/plugin"),
|
||||
)
|
||||
with patch("services.datasource_provider_service.get_current_user", return_value=mock_user):
|
||||
with pytest.raises(ValueError, match="already exists"):
|
||||
service.add_datasource_api_key_provider("clash", "t1", make_id(), {"sk": "v"})
|
||||
|
||||
def test_should_raise_value_error_when_credentials_validation_fails(self, service, mock_db_session, mock_user):
|
||||
mock_db_session.scalar.return_value = 0
|
||||
def test_should_raise_value_error_when_credentials_validation_fails(self, service, sqlite_session, mock_user):
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service.provider_manager, "validate_provider_credentials", side_effect=Exception("bad cred")),
|
||||
@ -599,18 +746,19 @@ class TestDatasourceProviderService:
|
||||
with pytest.raises(ValueError, match="Failed to validate"):
|
||||
service.add_datasource_api_key_provider("nm", "t1", make_id(), {"k": "v"})
|
||||
|
||||
def test_should_add_api_key_provider_and_commit_when_valid(self, service, mock_db_session, mock_user):
|
||||
mock_db_session.scalar.return_value = 0
|
||||
def test_should_add_api_key_provider_and_commit_when_valid(self, service, sqlite_session, mock_user):
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service.provider_manager, "validate_provider_credentials"),
|
||||
patch.object(service, "extract_secret_variables", return_value=["sk"]),
|
||||
):
|
||||
service.add_datasource_api_key_provider(None, "t1", make_id(), {"sk": "v"})
|
||||
mock_db_session.add.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
provider = sqlite_session.scalar(select(DatasourceProvider))
|
||||
assert provider is not None
|
||||
assert provider.encrypted_credentials == {"sk": "enc_tok"}
|
||||
|
||||
def test_should_acquire_redis_lock_when_adding_api_key_provider(self, service, mock_db_session, mock_user):
|
||||
mock_db_session.scalar.return_value = 0
|
||||
def test_should_acquire_redis_lock_when_adding_api_key_provider(self, service, sqlite_session, mock_user):
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service.provider_manager, "validate_provider_credentials"),
|
||||
@ -655,25 +803,22 @@ class TestDatasourceProviderService:
|
||||
# list_datasource_credentials (lines 721-754)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_empty_list_when_no_credentials_stored(self, service, mock_db_session):
|
||||
mock_db_session.scalars.return_value.all.return_value = []
|
||||
assert service.list_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) == []
|
||||
def test_should_return_empty_list_when_no_credentials_stored(self, service, sqlite_session):
|
||||
assert service.list_datasource_credentials("t1", "prov", "org/plug", session=sqlite_session) == []
|
||||
|
||||
def test_should_return_masked_credentials_list_when_credentials_exist(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "api_key"
|
||||
p.encrypted_credentials = {"sk": "v"}
|
||||
p.is_default = False
|
||||
mock_db_session.scalars.return_value.all.return_value = [p]
|
||||
def test_should_return_masked_credentials_list_when_credentials_exist(self, service, sqlite_session):
|
||||
p = make_provider(encrypted_credentials={"sk": "v"})
|
||||
persist(sqlite_session, p)
|
||||
with patch.object(service, "extract_secret_variables", return_value=["sk"]):
|
||||
result = service.list_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session)
|
||||
result = service.list_datasource_credentials("t1", "prov", "org/plug", session=sqlite_session)
|
||||
assert len(result) == 1
|
||||
assert result[0]["credential"] == {"sk": "obf"}
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# get_all_datasource_credentials (lines 808-871)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_aggregate_credentials_for_non_hardcoded_plugin(self, service):
|
||||
def test_should_aggregate_credentials_for_non_hardcoded_plugin(self, service, sqlite_session):
|
||||
with patch("services.datasource_provider_service.PluginDatasourceManager") as mock_mgr:
|
||||
ds = MagicMock()
|
||||
ds.provider = "prov"
|
||||
@ -682,12 +827,10 @@ class TestDatasourceProviderService:
|
||||
mock_mgr.return_value.fetch_installed_datasource_providers.return_value = [ds]
|
||||
cred = {"credential": {"k": "v"}, "is_default": True}
|
||||
with patch.object(service, "list_datasource_credentials", return_value=[cred]):
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = 0
|
||||
results = service.get_all_datasource_credentials("t1", session=session)
|
||||
results = service.get_all_datasource_credentials("t1", session=sqlite_session)
|
||||
assert len(results) == 1
|
||||
|
||||
def test_should_include_oauth_schema_for_hardcoded_plugin_ids(self, service, mock_db_session):
|
||||
def test_should_include_oauth_schema_for_hardcoded_plugin_ids(self, service, sqlite_session):
|
||||
"""Lines 819-871: get_all_datasource_credentials covers hardcoded langgenius plugin IDs."""
|
||||
with patch("services.datasource_provider_service.PluginDatasourceManager") as mock_mgr:
|
||||
ds = MagicMock()
|
||||
@ -709,7 +852,7 @@ class TestDatasourceProviderService:
|
||||
patch.object(service, "is_tenant_oauth_params_enabled", return_value=False),
|
||||
patch.object(service, "is_system_oauth_params_exist", return_value=False),
|
||||
):
|
||||
results = service.get_all_datasource_credentials("t1", session=mock_db_session)
|
||||
results = service.get_all_datasource_credentials("t1", session=sqlite_session)
|
||||
assert len(results) == 1
|
||||
assert results[0]["oauth_schema"] is not None
|
||||
|
||||
@ -717,47 +860,39 @@ class TestDatasourceProviderService:
|
||||
# get_real_datasource_credentials (lines 873-915)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_empty_list_when_no_real_credentials_exist(self, service, mock_db_session):
|
||||
mock_db_session.scalars.return_value.all.return_value = []
|
||||
assert service.get_real_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) == []
|
||||
def test_should_return_empty_list_when_no_real_credentials_exist(self, service, sqlite_session):
|
||||
assert service.get_real_datasource_credentials("t1", "prov", "org/plug", session=sqlite_session) == []
|
||||
|
||||
def test_should_return_decrypted_credential_list_when_credentials_exist(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.auth_type = "api_key"
|
||||
p.encrypted_credentials = {"sk": "v"}
|
||||
mock_db_session.scalars.return_value.all.return_value = [p]
|
||||
def test_should_return_decrypted_credential_list_when_credentials_exist(self, service, sqlite_session):
|
||||
p = make_provider(encrypted_credentials={"sk": "v"})
|
||||
persist(sqlite_session, p)
|
||||
with patch.object(service, "extract_secret_variables", return_value=["sk"]):
|
||||
result = service.get_real_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session)
|
||||
result = service.get_real_datasource_credentials("t1", "prov", "org/plug", session=sqlite_session)
|
||||
assert len(result) == 1
|
||||
assert result[0]["credentials"] == {"sk": "dec_tok"}
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# update_datasource_credentials (lines 917-978)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_raise_value_error_when_credential_not_found_on_update(self, service, mock_db_session, mock_user):
|
||||
mock_db_session.scalar.return_value = None
|
||||
def test_should_raise_value_error_when_credential_not_found_on_update(self, service, sqlite_session, mock_user):
|
||||
with patch("services.datasource_provider_service.get_current_user", return_value=mock_user):
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
service.update_datasource_credentials("t1", "id", "prov", "org/plug", {}, "name")
|
||||
|
||||
def test_should_raise_value_error_when_new_name_already_used_on_update(self, service, mock_db_session, mock_user):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.name = "old_name"
|
||||
p.auth_type = "api_key"
|
||||
p.encrypted_credentials = {"sk": "e"}
|
||||
mock_db_session.scalar.side_effect = [p, 1] # first: fetch provider, second: name conflict count
|
||||
def test_should_raise_value_error_when_new_name_already_used_on_update(self, service, sqlite_session, mock_user):
|
||||
p = make_provider(credential_id="id", name="old_name", encrypted_credentials={"sk": "e"})
|
||||
conflict = make_provider(credential_id="conflict-id", name="new_name")
|
||||
persist(sqlite_session, p, conflict)
|
||||
with patch("services.datasource_provider_service.get_current_user", return_value=mock_user):
|
||||
with pytest.raises(ValueError, match="already exists"):
|
||||
service.update_datasource_credentials("t1", "id", "prov", "org/plug", {}, "new_name")
|
||||
|
||||
def test_should_raise_value_error_when_credential_validation_fails_on_update(
|
||||
self, service, mock_db_session, mock_user
|
||||
self, service, sqlite_session, mock_user
|
||||
):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.name = "old_name"
|
||||
p.auth_type = "api_key"
|
||||
p.encrypted_credentials = {"sk": "e"}
|
||||
mock_db_session.scalar.side_effect = [p, 0] # first: fetch provider, second: name conflict count
|
||||
p = make_provider(credential_id="id", name="old_name", encrypted_credentials={"sk": "e"})
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service, "extract_secret_variables", return_value=["sk"]),
|
||||
@ -766,35 +901,33 @@ class TestDatasourceProviderService:
|
||||
with pytest.raises(ValueError, match="Failed to validate"):
|
||||
service.update_datasource_credentials("t1", "id", "prov", "org/plug", {"sk": "v"}, "name")
|
||||
|
||||
def test_should_encrypt_credentials_and_commit_when_update_succeeds(self, service, mock_db_session, mock_user):
|
||||
def test_should_encrypt_credentials_and_commit_when_update_succeeds(self, service, sqlite_session, mock_user):
|
||||
"""Verifies that encrypted_credentials is reassigned with encrypted value and commit is called."""
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
p.name = "old_name"
|
||||
p.auth_type = "api_key"
|
||||
p.encrypted_credentials = {"sk": "old_enc"}
|
||||
mock_db_session.scalar.side_effect = [p, 0] # first: fetch provider, second: name conflict count
|
||||
p = make_provider(credential_id="id", name="old_name", encrypted_credentials={"sk": "old_enc"})
|
||||
persist(sqlite_session, p)
|
||||
with (
|
||||
patch("services.datasource_provider_service.get_current_user", return_value=mock_user),
|
||||
patch.object(service, "extract_secret_variables", return_value=["sk"]),
|
||||
patch.object(service.provider_manager, "validate_provider_credentials"),
|
||||
):
|
||||
service.update_datasource_credentials("t1", "id", "prov", "org/plug", {"sk": "new_val"}, "name")
|
||||
# encrypter must have been called with the new secret value
|
||||
self._enc.encrypt_token.assert_called()
|
||||
# commit must be called exactly once
|
||||
sqlite_session.expire_all()
|
||||
updated = sqlite_session.get(DatasourceProvider, p.id)
|
||||
assert updated.name == "name"
|
||||
assert updated.encrypted_credentials == {"sk": "enc_tok"}
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# remove_datasource_credentials (lines 980-997)
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_delete_provider_and_commit_when_found(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
mock_db_session.scalar.return_value = p
|
||||
service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=mock_db_session)
|
||||
mock_db_session.delete.assert_called_once_with(p)
|
||||
def test_should_delete_provider_and_commit_when_found(self, service, sqlite_session):
|
||||
p = make_provider(credential_id="id")
|
||||
persist(sqlite_session, p)
|
||||
service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=sqlite_session)
|
||||
assert sqlite_session.get(DatasourceProvider, p.id) is None
|
||||
|
||||
def test_should_do_nothing_when_credential_not_found_on_remove(self, service, mock_db_session):
|
||||
def test_should_do_nothing_when_credential_not_found_on_remove(self, service, sqlite_session):
|
||||
"""No error raised; no delete called when record doesn't exist (lines 994 branch)."""
|
||||
mock_db_session.scalar.return_value = None
|
||||
service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=mock_db_session)
|
||||
mock_db_session.delete.assert_not_called()
|
||||
service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=sqlite_session)
|
||||
assert sqlite_session.scalars(select(DatasourceProvider)).all() == []
|
||||
|
||||
Loading…
Reference in New Issue
Block a user