test: use sqlite3 session in test_datasource_provider_service (#38695)

This commit is contained in:
Asuka Minato 2026-07-25 12:16:43 +09:00 committed by GitHub
parent a2d9aeff37
commit 8906a49e56
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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() == []