SIENTIAPDE-1171

Refactor model_repository and enhance test coverage for MLFlow functionalities

- Updated model_repository to ensure the 'temp' directory is created if it doesn't exist using `exist_ok=True`.
- Added new tests for retraining and updating production models, including error handling scenarios.
- Improved existing tests for model management workflows to ensure robustness and reliability.
This commit is contained in:
vitor-aignosi
2025-07-24 10:24:08 -03:00
parent 4283730e7a
commit c738e8b0df
4 changed files with 280 additions and 12 deletions

View File

@@ -164,6 +164,18 @@ def test_get_experiment_last_run(mlflow, mlflow_repository):
assert output == '2'
@patch('laborious.utils.repository.model_repository.mlflow')
def test_get_experiment_last_run_error(mlflow, mlflow_repository):
mlflow.search_runs.return_value = []
try:
mlflow_repository.get_experiment_last_run(0)
except ValueError as e:
assert str(e) == 'Runs is not a pandas DataFrame'
else:
assert False
@patch('laborious.utils.repository.model_repository.mlflow.sklearn')
@patch('laborious.utils.repository.model_repository.mlflow.set_experiment')
def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
@@ -233,14 +245,6 @@ def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
@patch('laborious.utils.repository.model_repository.mlflow.log_artifact')
def test_perform_model_retrain(log_artifact, log_model, log_param, start_run, mlflow_repository):
_os_remove = patch(
'laborious.utils.repository.model_repository.remove')
_os_makedirs = patch(
'laborious.utils.repository.model_repository.makedirs')
_os_path_exists = patch(
'laborious.utils.repository.model_repository.path.exists',
MagicMock(return_value=False))
prediction_model_mock = MagicMock()
data_model_mock = MagicMock()
experiment = 'test'
@@ -277,6 +281,27 @@ def test_perform_model_retrain(log_artifact, log_model, log_param, start_run, ml
assert output == ("Model retrained successfully", experiment)
def test_retrain_model(mlflow_repository):
data = MagicMock()
model_name = 'test'
mlflow_repository.create_model_experiment = MagicMock(
return_value=('data_model', 'prediction_model', '0'))
mlflow_repository.perform_model_retrain = MagicMock(
return_value='Model retrained successfully')
output = mlflow_repository.retrain_model(data, model_name)
mlflow_repository.create_model_experiment.assert_called_once_with(
model_name, data)
mlflow_repository.perform_model_retrain.assert_called_once_with(
'data_model', 'prediction_model', '0', model_name, data)
assert output == 'Model retrained successfully'
@patch('laborious.utils.repository.model_repository.mlflow')
def test_update_production_model_by_run_id(mlflow, mlflow_repository):
client_mock = MagicMock()
@@ -312,6 +337,24 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
}
@patch('laborious.utils.repository.model_repository.mlflow')
def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
mlflow.tracking.MlflowClient.return_value = MagicMock(
get_registered_model=MagicMock(
return_value=MagicMock(
latest_versions={}
)
)
)
try:
mlflow_repository.update_production_model_by_run_id('0', 'test')
except Exception as e:
assert str(e) == 'Model versions is not a list'
else:
assert False
def test_update_production_model(mlflow_repository):
connector = mlflow_repository