mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
test: use sqlite3 session in test_provider_configuration (#38689)
This commit is contained in:
parent
3afd81a028
commit
15acaf1fa6
@ -1,7 +1,12 @@
|
|||||||
|
"""Provider-configuration behavior with persisted SQLite credential records."""
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import Mock, patch
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from core.entities import provider_configuration as provider_configuration_module
|
||||||
from core.entities.provider_configuration import ProviderConfiguration, SystemConfigurationStatus
|
from core.entities.provider_configuration import ProviderConfiguration, SystemConfigurationStatus
|
||||||
from core.entities.provider_entities import (
|
from core.entities.provider_entities import (
|
||||||
CustomConfiguration,
|
CustomConfiguration,
|
||||||
@ -21,7 +26,7 @@ from graphon.model_runtime.entities.provider_entities import (
|
|||||||
FormType,
|
FormType,
|
||||||
ProviderEntity,
|
ProviderEntity,
|
||||||
)
|
)
|
||||||
from models.provider import Provider, ProviderType
|
from models.provider import Provider, ProviderCredential, ProviderType
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
@ -207,33 +212,25 @@ class TestProviderConfiguration:
|
|||||||
# Assert
|
# Assert
|
||||||
assert result is False
|
assert result is False
|
||||||
|
|
||||||
@patch("core.entities.provider_configuration.Session")
|
@pytest.mark.parametrize("sqlite_session", [(Provider,)], indirect=True)
|
||||||
def test_get_provider_record_found(self, mock_session, provider_configuration):
|
def test_get_provider_record_found(self, provider_configuration, sqlite_session: Session):
|
||||||
"""Test getting provider record successfully"""
|
"""Test getting provider record successfully"""
|
||||||
# Arrange
|
provider = Provider(tenant_id="test_tenant", provider_name="openai")
|
||||||
mock_provider = Mock(spec=Provider)
|
sqlite_session.add(provider)
|
||||||
mock_session_instance = Mock()
|
sqlite_session.commit()
|
||||||
mock_session.return_value.__enter__.return_value = mock_session_instance
|
|
||||||
mock_session_instance.execute.return_value.scalar_one_or_none.return_value = mock_provider
|
|
||||||
|
|
||||||
# Act
|
result = provider_configuration._get_provider_record(sqlite_session)
|
||||||
result = provider_configuration._get_provider_record(mock_session_instance)
|
|
||||||
|
|
||||||
# Assert
|
assert result is provider
|
||||||
assert result == mock_provider
|
|
||||||
|
|
||||||
@patch("core.entities.provider_configuration.Session")
|
@pytest.mark.parametrize("sqlite_session", [(Provider,)], indirect=True)
|
||||||
def test_get_provider_record_not_found(self, mock_session, provider_configuration):
|
def test_get_provider_record_not_found(self, provider_configuration, sqlite_session: Session):
|
||||||
"""Test getting provider record when not found"""
|
"""Test getting provider record when not found"""
|
||||||
# Arrange
|
sqlite_session.add(Provider(tenant_id="other_tenant", provider_name="openai"))
|
||||||
mock_session_instance = Mock()
|
sqlite_session.commit()
|
||||||
mock_session.return_value.__enter__.return_value = mock_session_instance
|
|
||||||
mock_session_instance.execute.return_value.scalar_one_or_none.return_value = None
|
|
||||||
|
|
||||||
# Act
|
result = provider_configuration._get_provider_record(sqlite_session)
|
||||||
result = provider_configuration._get_provider_record(mock_session_instance)
|
|
||||||
|
|
||||||
# Assert
|
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
def test_init_with_customizable_model_only(
|
def test_init_with_customizable_model_only(
|
||||||
@ -270,47 +267,47 @@ class TestProviderConfiguration:
|
|||||||
assert credentials is not None
|
assert credentials is not None
|
||||||
assert "openai_api_key" in credentials
|
assert "openai_api_key" in credentials
|
||||||
|
|
||||||
@patch("core.entities.provider_configuration.Session")
|
@pytest.mark.parametrize("sqlite_session", [(Provider, ProviderCredential)], indirect=True)
|
||||||
def test_get_specific_provider_credential_success(self, mock_session, provider_configuration):
|
def test_get_specific_provider_credential_success(
|
||||||
|
self, monkeypatch: pytest.MonkeyPatch, provider_configuration, sqlite_session: Session
|
||||||
|
):
|
||||||
"""Test getting specific provider credential successfully"""
|
"""Test getting specific provider credential successfully"""
|
||||||
# Arrange
|
|
||||||
credential_id = "test_credential_id"
|
credential_id = "test_credential_id"
|
||||||
mock_credential = Mock()
|
credential = ProviderCredential(
|
||||||
mock_credential.encrypted_config = '{"openai_api_key": "encrypted_key"}'
|
tenant_id="test_tenant",
|
||||||
|
provider_name="openai",
|
||||||
|
credential_name="primary",
|
||||||
|
encrypted_config='{"openai_api_key": "encrypted_key"}',
|
||||||
|
)
|
||||||
|
credential.id = credential_id
|
||||||
|
sqlite_session.add_all((Provider(tenant_id="test_tenant", provider_name="openai"), credential))
|
||||||
|
sqlite_session.commit()
|
||||||
|
monkeypatch.setattr(
|
||||||
|
provider_configuration_module,
|
||||||
|
"db",
|
||||||
|
SimpleNamespace(engine=sqlite_session.get_bind()),
|
||||||
|
)
|
||||||
|
|
||||||
mock_session_instance = Mock()
|
result = provider_configuration._get_specific_provider_credential(credential_id)
|
||||||
mock_session.return_value.__enter__.return_value = mock_session_instance
|
|
||||||
mock_session_instance.execute.return_value.scalar_one_or_none.return_value = mock_credential
|
|
||||||
|
|
||||||
# Act
|
assert result == {"openai_api_key": "encrypted_key"}
|
||||||
with patch.object(provider_configuration, "_get_specific_provider_credential") as mock_get:
|
|
||||||
mock_get.return_value = {"openai_api_key": "test_key"}
|
|
||||||
result = provider_configuration._get_specific_provider_credential(credential_id)
|
|
||||||
|
|
||||||
# Assert
|
@pytest.mark.parametrize("sqlite_session", [(Provider, ProviderCredential)], indirect=True)
|
||||||
assert result == {"openai_api_key": "test_key"}
|
def test_get_specific_provider_credential_not_found(
|
||||||
|
self, monkeypatch: pytest.MonkeyPatch, provider_configuration, sqlite_session: Session
|
||||||
@patch("core.entities.provider_configuration.Session")
|
):
|
||||||
def test_get_specific_provider_credential_not_found(self, mock_session, provider_configuration):
|
|
||||||
"""Test getting specific provider credential when not found"""
|
"""Test getting specific provider credential when not found"""
|
||||||
# Arrange
|
|
||||||
credential_id = "nonexistent_credential_id"
|
credential_id = "nonexistent_credential_id"
|
||||||
|
sqlite_session.add(Provider(tenant_id="other_tenant", provider_name="openai"))
|
||||||
|
sqlite_session.commit()
|
||||||
|
monkeypatch.setattr(
|
||||||
|
provider_configuration_module,
|
||||||
|
"db",
|
||||||
|
SimpleNamespace(engine=sqlite_session.get_bind()),
|
||||||
|
)
|
||||||
|
|
||||||
mock_session_instance = Mock()
|
with pytest.raises(ValueError, match=credential_id):
|
||||||
mock_session.return_value.__enter__.return_value = mock_session_instance
|
provider_configuration._get_specific_provider_credential(credential_id)
|
||||||
mock_session_instance.execute.return_value.scalar_one_or_none.return_value = None
|
|
||||||
|
|
||||||
# Act & Assert
|
|
||||||
with patch.object(provider_configuration, "_get_specific_provider_credential") as mock_get:
|
|
||||||
mock_get.return_value = None
|
|
||||||
result = provider_configuration._get_specific_provider_credential(credential_id)
|
|
||||||
assert result is None
|
|
||||||
|
|
||||||
# Act
|
|
||||||
credentials = provider_configuration.get_current_credentials(ModelType.LLM, "gpt-4")
|
|
||||||
|
|
||||||
# Assert
|
|
||||||
assert credentials == {"openai_api_key": "test_key"}
|
|
||||||
|
|
||||||
def test_extract_secret_variables_with_secret_input(self, provider_configuration):
|
def test_extract_secret_variables_with_secret_input(self, provider_configuration):
|
||||||
"""Test extracting secret variables from credential form schemas"""
|
"""Test extracting secret variables from credential form schemas"""
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user