"""Tests for model utils.""" import pytest from territories import config from territories.connectors import mysql from territories.connectors import sqlite from territories.util import model_util @pytest.mark.parametrize( 'feature_flag_value, expected_base_model, environment', [(0, sqlite.BaseModel, config.TEST_ENVIRONMENT), (1, mysql.BaseModel, config.DEV_ENVIRONMENT), (1, sqlite.BaseModel, config.TEST_ENVIRONMENT)]) def test_get_base_model( mocker, feature_flag_value, expected_base_model, environment): """Expect to get BaseModel according to feature flag value.""" mock_config = mocker.patch('territories.util.model_util.config') mock_config.RDS_OWS_TERRITORIES = feature_flag_value mock_config.ENVIRONMENT = environment mock_config.TEST_ENVIRONMENT = config.TEST_ENVIRONMENT mock_config.DEV_ENVIRONMENT = config.DEV_ENVIRONMENT base_model = model_util.get_base_model() assert base_model == expected_base_model @pytest.mark.parametrize( 'feature_flag_value, expected_session_scope', [(0, sqlite.session_scope), (1, mysql.db_session)]) def test_get_session_scope( mocker, feature_flag_value, expected_session_scope): """Expect to get session_scope according to feature flag value.""" mock_config = mocker.patch('territories.util.model_util.config') mock_config.RDS_OWS_TERRITORIES = feature_flag_value session_scope = model_util.get_session_scope() assert session_scope == expected_session_scope