From 272e02dadcb186132f9fbcd189f1080430c7b342 Mon Sep 17 00:00:00 2001 From: vitor-aignosi Date: Thu, 11 Jun 2026 13:34:25 -0300 Subject: [PATCH] SIENTIAPDE-1646 SIENTIAPDE-1646 Refactor MLFlow tests and update artifact handling - Enhanced test cases for MLFlow to improve clarity and accuracy in data handling. - Updated references in tests to use 'evaluation_data.csv' and 'test_data.csv' instead of 'retrain_input.csv' and 'train_data.csv'. - Introduced a new helper function for creating prediction frames to streamline test setup. --- laborious/activities/mlflow.py | 2 +- laborious/worker/worker.py | 4 +- tests/laborious/activities/test_mlflow.py | 60 ++++++++++++++--------- 3 files changed, 40 insertions(+), 26 deletions(-) diff --git a/laborious/activities/mlflow.py b/laborious/activities/mlflow.py index 3ac5d95..bbc3e45 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -559,7 +559,7 @@ class MLFlow(SientiaMonitoring): evaluation_csv = Path(tmp_dir) / 'evaluation_data.csv' data.to_csv(raw_csv, index=False) evaluation_data.to_csv(evaluation_csv, index=False) - + mlflow.log_artifact(str(raw_csv)) mlflow.log_artifact(str(evaluation_csv)) finally: diff --git a/laborious/worker/worker.py b/laborious/worker/worker.py index d5f3c9a..f284a8c 100644 --- a/laborious/worker/worker.py +++ b/laborious/worker/worker.py @@ -216,7 +216,7 @@ async def main(): activities.export_data_to_postgres, ], logger=logger, - runtime='core' + runtime='core', ), prepare_worker( temporal_client=temporal_client, @@ -229,7 +229,7 @@ async def main(): activities.export_data_to_postgres, ], logger=logger, - runtime='core' + runtime='core', ), prepare_worker( temporal_client=temporal_client, diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py index 030ad16..43dd0ae 100644 --- a/tests/laborious/activities/test_mlflow.py +++ b/tests/laborious/activities/test_mlflow.py @@ -399,6 +399,8 @@ def test_retrain_model_success_data_success_retrain( 'value': [1.0, 2.0], } ) + pivoted_index = pd.DatetimeIndex([ts]) + wrapper.retrain.return_value = _retrain_prediction_frame(pivoted_index, 'target', [1.0]) payload = MagicMock() payload.retrieve = MagicMock(return_value=raw_data) @@ -428,6 +430,7 @@ def test_retrain_model_success_with_payload_data( mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow ): mock_mkdtemp.return_value = 'tmp' + mock_to_datetime.side_effect = lambda idx, **kwargs: pd.DatetimeIndex(idx) mv_alias = MagicMock(run_id='src') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias wrapper = MagicMock() @@ -437,20 +440,21 @@ def test_retrain_model_success_with_payload_data( mock_cm.__exit__.return_value = False mlflow.mlflow_repository.start_run.return_value = mock_cm - raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at']) - raw_data.__getitem__.return_value.max.return_value = 'ts' + ts = pd.Timestamp('2020-01-01', tz='UTC') + raw_data = pd.DataFrame( + { + 'variable': ['target', 'f1', 'target', 'f1'], + 'timestamp': [ts, ts, ts + pd.Timedelta(hours=1), ts + pd.Timedelta(hours=1)], + 'value': [1.0, 2.0, 3.0, 4.0], + 'created_at': [ts, ts, ts + pd.Timedelta(hours=1), ts + pd.Timedelta(hours=1)], + } + ) + pivoted_index = pd.DatetimeIndex([ts, ts + pd.Timedelta(hours=1)]) + wrapper.retrain.return_value = _retrain_prediction_frame(pivoted_index, 'target', [1.0, 3.0]) + payload = MagicMock() payload.retrieve = MagicMock(return_value=raw_data) - pivoted = MagicMock() - raw_data.sort_values.return_value = raw_data - raw_data.drop_duplicates.return_value = raw_data - raw_data.pivot.return_value = pivoted - pivoted.fillna = MagicMock() - pivoted.columns.name = None - pivoted.index = MagicMock() - pivoted.__setitem__ = MagicMock() - response = mlflow.retrain_model( { **metadata, @@ -482,13 +486,17 @@ def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag( mlflow.mlflow_repository.start_run.return_value = mock_cm ts = pd.Timestamp('2020-01-01', tz='UTC') + ts_next = ts + pd.Timedelta(days=1) raw_data = pd.DataFrame( { 'variable': ['target', 'f1', 'target', 'f1'], - 'timestamp': [ts, ts, ts + pd.Timedelta(days=1), ts + pd.Timedelta(days=1)], + 'timestamp': [ts, ts, ts_next, ts_next], 'value': [1.0, 2.0, 3.0, 4.0], } ) + pivoted_index = pd.DatetimeIndex([ts, ts_next]) + wrapper.retrain.return_value = _retrain_prediction_frame(pivoted_index, 'target', [1.0, 3.0]) + payload = MagicMock() payload.retrieve = MagicMock(return_value=raw_data) @@ -658,6 +666,12 @@ def _artifact_file_info(path: str) -> MagicMock: return file_info +def _retrain_prediction_frame( + index: pd.DatetimeIndex, target: str, values: list[float] +) -> pd.DataFrame: + return pd.DataFrame({target: values}, index=index) + + @patch('laborious.activities.mlflow.to_datetime') def test_get_reference_data_success(mock_to_datetime, mlflow): input_data = { @@ -668,7 +682,7 @@ def test_get_reference_data_success(mock_to_datetime, mlflow): mv = MagicMock(run_id='run1') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv mlflow.mlflow_repository._client.list_artifacts.return_value = [ - _artifact_file_info('retrain_input.csv'), + _artifact_file_info('evaluation_data.csv'), ] mock_reference_data = MagicMock() @@ -690,7 +704,7 @@ def test_get_reference_data_success(mock_to_datetime, mlflow): mlflow.mlflow_repository.download_artifacts.assert_called_once_with( run_id='run1', - artifact_path='retrain_input.csv', + artifact_path='evaluation_data.csv', dst_path='/t', metadata=metadata['metadata'], ) @@ -712,7 +726,7 @@ def test_get_reference_data_not_found(mlflow): assert result is None -def test_get_reference_data_only_train_data_csv(mlflow): +def test_get_reference_data_only_test_data_csv(mlflow): input_data = { **metadata, 'model_name': 'test_model', @@ -720,7 +734,7 @@ def test_get_reference_data_only_train_data_csv(mlflow): mv = MagicMock(run_id='run1') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv mlflow.mlflow_repository._client.list_artifacts.return_value = [ - _artifact_file_info('train_data.csv'), + _artifact_file_info('test_data.csv'), ] mock_reference_data = MagicMock() @@ -741,14 +755,14 @@ def test_get_reference_data_only_train_data_csv(mlflow): mlflow.mlflow_repository.download_artifacts.assert_called_once_with( run_id='run1', - artifact_path='train_data.csv', + artifact_path='test_data.csv', dst_path='/t', metadata=metadata['metadata'], ) assert result == mock_reference_data.to_dict.return_value -def test_get_reference_data_prefers_retrain_input_when_both_listed(mlflow): +def test_get_reference_data_prefers_evaluation_data_when_both_listed(mlflow): input_data = { **metadata, 'model_name': 'test_model', @@ -756,8 +770,8 @@ def test_get_reference_data_prefers_retrain_input_when_both_listed(mlflow): mv = MagicMock(run_id='run1') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv mlflow.mlflow_repository._client.list_artifacts.return_value = [ - _artifact_file_info('train_data.csv'), - _artifact_file_info('retrain_input.csv'), + _artifact_file_info('test_data.csv'), + _artifact_file_info('evaluation_data.csv'), ] mock_reference_data = MagicMock() @@ -778,7 +792,7 @@ def test_get_reference_data_prefers_retrain_input_when_both_listed(mlflow): mlflow.mlflow_repository.download_artifacts.assert_called_once_with( run_id='run1', - artifact_path='retrain_input.csv', + artifact_path='evaluation_data.csv', dst_path='/t', metadata=metadata['metadata'], ) @@ -811,7 +825,7 @@ def test_get_reference_data_missing_csv_file_returns_none(mlflow): mv = MagicMock(run_id='run1') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv mlflow.mlflow_repository._client.list_artifacts.return_value = [ - _artifact_file_info('retrain_input.csv'), + _artifact_file_info('evaluation_data.csv'), ] with patch('laborious.activities.mlflow.tempfile.mkdtemp', return_value='tmp'): @@ -833,7 +847,7 @@ def test_get_reference_data_exception(mlflow): mv = MagicMock(run_id='run1') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv mlflow.mlflow_repository._client.list_artifacts.return_value = [ - _artifact_file_info('retrain_input.csv'), + _artifact_file_info('evaluation_data.csv'), ] mlflow.mlflow_repository.download_artifacts.side_effect = Exception('dl fail')