From 8906a49e56557cb2488409bdfec2eeafdf6937cc Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Sat, 25 Jul 2026 12:16:43 +0900 Subject: [PATCH] test: use sqlite3 session in test_datasource_provider_service (#38695) --- .../test_datasource_provider_service.py | 631 +++++++++++------- 1 file changed, 382 insertions(+), 249 deletions(-) diff --git a/api/tests/unit_tests/services/test_datasource_provider_service.py b/api/tests/unit_tests/services/test_datasource_provider_service.py index bd6891d846a..4db672f393c 100644 --- a/api/tests/unit_tests/services/test_datasource_provider_service.py +++ b/api/tests/unit_tests/services/test_datasource_provider_service.py @@ -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() == []