mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: select in datasource_provider_service (#34548)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
@@ -57,6 +57,10 @@ class TestDatasourceProviderService:
|
||||
q.count.return_value = 0
|
||||
q.delete.return_value = 1
|
||||
|
||||
# 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
|
||||
|
||||
@@ -183,11 +187,11 @@ class TestDatasourceProviderService:
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_true_when_tenant_oauth_params_enabled(self, service, mock_db_session):
|
||||
mock_db_session.query().count.return_value = 1
|
||||
mock_db_session.scalar.return_value = 1
|
||||
assert service.is_tenant_oauth_params_enabled("t1", make_id()) is True
|
||||
|
||||
def test_should_return_false_when_tenant_oauth_params_disabled(self, service, mock_db_session):
|
||||
mock_db_session.query().count.return_value = 0
|
||||
mock_db_session.scalar.return_value = 0
|
||||
assert service.is_tenant_oauth_params_enabled("t1", make_id()) is False
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
@@ -401,7 +405,7 @@ class TestDatasourceProviderService:
|
||||
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.query().first.return_value = tenant_params
|
||||
mock_db_session.scalar.return_value = 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)
|
||||
assert result == {"k": "mask"}
|
||||
@@ -409,13 +413,13 @@ class TestDatasourceProviderService:
|
||||
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.query().first.return_value = tenant_params
|
||||
mock_db_session.scalar.return_value = 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)
|
||||
assert result == {"k": "dec"}
|
||||
|
||||
def test_should_return_none_when_no_tenant_oauth_config_exists(self, service, mock_db_session):
|
||||
mock_db_session.query().first.return_value = None
|
||||
mock_db_session.scalar.return_value = None
|
||||
assert service.get_tenant_oauth_client("t1", make_id()) is None
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
@@ -616,7 +620,7 @@ class TestDatasourceProviderService:
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_empty_list_when_no_credentials_stored(self, service, mock_db_session):
|
||||
mock_db_session.query().all.return_value = []
|
||||
mock_db_session.scalars.return_value.all.return_value = []
|
||||
assert service.list_datasource_credentials("t1", "prov", "org/plug") == []
|
||||
|
||||
def test_should_return_masked_credentials_list_when_credentials_exist(self, service, mock_db_session):
|
||||
@@ -624,7 +628,7 @@ class TestDatasourceProviderService:
|
||||
p.auth_type = "api_key"
|
||||
p.encrypted_credentials = {"sk": "v"}
|
||||
p.is_default = False
|
||||
mock_db_session.query().all.return_value = [p]
|
||||
mock_db_session.scalars.return_value.all.return_value = [p]
|
||||
with patch.object(service, "extract_secret_variables", return_value=["sk"]):
|
||||
result = service.list_datasource_credentials("t1", "prov", "org/plug")
|
||||
assert len(result) == 1
|
||||
@@ -676,14 +680,14 @@ class TestDatasourceProviderService:
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_should_return_empty_list_when_no_real_credentials_exist(self, service, mock_db_session):
|
||||
mock_db_session.query().all.return_value = []
|
||||
mock_db_session.scalars.return_value.all.return_value = []
|
||||
assert service.get_real_datasource_credentials("t1", "prov", "org/plug") == []
|
||||
|
||||
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.query().all.return_value = [p]
|
||||
mock_db_session.scalars.return_value.all.return_value = [p]
|
||||
with patch.object(service, "extract_secret_variables", return_value=["sk"]):
|
||||
result = service.get_real_datasource_credentials("t1", "prov", "org/plug")
|
||||
assert len(result) == 1
|
||||
@@ -751,13 +755,13 @@ class TestDatasourceProviderService:
|
||||
|
||||
def test_should_delete_provider_and_commit_when_found(self, service, mock_db_session):
|
||||
p = MagicMock(spec=DatasourceProvider)
|
||||
mock_db_session.query().first.return_value = p
|
||||
mock_db_session.scalar.return_value = p
|
||||
service.remove_datasource_credentials("t1", "id", "prov", "org/plug")
|
||||
mock_db_session.delete.assert_called_once_with(p)
|
||||
mock_db_session.commit.assert_called_once()
|
||||
|
||||
def test_should_do_nothing_when_credential_not_found_on_remove(self, service, mock_db_session):
|
||||
"""No error raised; no delete called when record doesn't exist (lines 994 branch)."""
|
||||
mock_db_session.query().first.return_value = None
|
||||
mock_db_session.scalar.return_value = None
|
||||
service.remove_datasource_credentials("t1", "id", "prov", "org/plug")
|
||||
mock_db_session.delete.assert_not_called()
|
||||
|
||||
Reference in New Issue
Block a user