From 666188a7b30095ba3febc29e447128078a91ccaf Mon Sep 17 00:00:00 2001 From: PedroHMCosme Date: Wed, 19 Aug 2026 11:17:59 -0300 Subject: [PATCH] test(simple_metrics): rewrite tests against RegressionMetrics dispatch Rewrites 8 existing tests to mock RegressionMetrics instead of verifying manual numpy math. Adds 2 new tests: - r2 excluded + warning when model_type is non-linear - r2 included when model_type is absent (backward compat) Replaces silent-discard test with ValueError propagation test. SIENTIAPDE-1986 --- .../activities/test_model_metrics.py | 251 ++++++++++++------ 1 file changed, 168 insertions(+), 83 deletions(-) diff --git a/tests/laborious/activities/test_model_metrics.py b/tests/laborious/activities/test_model_metrics.py index 4c41109..ede6a27 100644 --- a/tests/laborious/activities/test_model_metrics.py +++ b/tests/laborious/activities/test_model_metrics.py @@ -45,6 +45,23 @@ metadata = { } +def _mock_regression_metrics_class(calculate_return): + """Returns a patch context manager that mocks RegressionMetrics.""" + mock_instance = MagicMock() + mock_instance.calculate.return_value = calculate_return + mock_class = MagicMock(return_value=mock_instance) + mock_class.is_r2_supported = MagicMock(return_value=True) + mock_class.supported_metrics = MagicMock(return_value=['rmse', 'mse', 'mae', 'r2']) + return ( + patch( + 'laborious.activities.model_metrics.RegressionMetrics', + mock_class, + ), + mock_class, + mock_instance, + ) + + def test_calculate_drift_invalid_chunk_period(model_metrics_activity): # Arrange input_data = { @@ -693,7 +710,6 @@ def test_get_drift_metrics_dataframe_error( def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity): - # Arrange target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'], @@ -710,28 +726,28 @@ def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity): 'interval_minutes': 5, } - # Act - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) - - # Assert - assert len(result['metric']) == 4 - assert 'rmse' in result['metric'].values - assert 'mse' in result['metric'].values - assert 'mae' in result['metric'].values - assert 'r2' in result['metric'].values - assert all(model_id == 'test_model_id' for model_id in result['model_id'].values) - assert all(timestamp == '2023-05-26 11:12:29' for timestamp in result['timestamp'].values) - assert all(data_size == 3 for data_size in result['data_size'].values) - assert all(interval_minutes == 5 for interval_minutes in result['interval_minutes'].values) - model_metrics_activity.info.assert_called_once_with( - "Calculating simple metrics for model test_model_id: ['rmse', 'mse', 'mae', 'r2']", - metadata['metadata'], + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [ + {'metric': 'rmse', 'value': 0.1}, + {'metric': 'mse', 'value': 0.01}, + {'metric': 'mae', 'value': 0.1}, + {'metric': 'r2', 'value': 0.99}, + ] ) - model_metrics_activity.debug.assert_called_once() + + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert len(result) == 4 + assert set(result['metric'].values) == {'rmse', 'mse', 'mae', 'r2'} + assert all(mid == 'test_model_id' for mid in result['model_id'].values) + assert all(ts == '2023-05-26 11:12:29' for ts in result['timestamp'].values) + assert all(ds == 3 for ds in result['data_size'].values) + assert all(im == 5 for im in result['interval_minutes'].values) + mock_instance.calculate.assert_called_once_with(['rmse', 'mse', 'mae', 'r2']) def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity): - # Arrange target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], @@ -748,23 +764,23 @@ def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity): 'interval_minutes': 5, } - # Act - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [{'metric': 'rmse', 'value': 0.1}] + ) - # Assert - assert len(result['metric']) == 1 + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert len(result) == 1 assert result['metric'].values[0] == 'rmse' assert result['model_id'].values[0] == 'test_model_id' assert result['timestamp'].values[0] == '2023-05-26 11:12:28' assert result['data_size'].values[0] == 2 assert result['interval_minutes'].values[0] == 5 - model_metrics_activity.info.assert_called_once_with( - "Calculating simple metrics for model test_model_id: ['rmse']", metadata['metadata'] - ) + mock_instance.calculate.assert_called_once_with(['rmse']) def test_calculate_simple_metrics_success_mse_only(model_metrics_activity): - # Arrange target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], @@ -781,23 +797,23 @@ def test_calculate_simple_metrics_success_mse_only(model_metrics_activity): 'interval_minutes': 5, } - # Act - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [{'metric': 'mse', 'value': 0.01}] + ) - # Assert - assert len(result['metric']) == 1 + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert len(result) == 1 assert result['metric'].values[0] == 'mse' assert result['model_id'].values[0] == 'test_model_id' assert result['timestamp'].values[0] == '2023-05-26 11:12:28' assert result['data_size'].values[0] == 2 assert result['interval_minutes'].values[0] == 5 - model_metrics_activity.info.assert_called_once_with( - "Calculating simple metrics for model test_model_id: ['mse']", metadata['metadata'] - ) + mock_instance.calculate.assert_called_once_with(['mse']) def test_calculate_simple_metrics_success_mae_only(model_metrics_activity): - # Arrange target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], @@ -814,23 +830,23 @@ def test_calculate_simple_metrics_success_mae_only(model_metrics_activity): 'interval_minutes': 5, } - # Act - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [{'metric': 'mae', 'value': 0.1}] + ) - # Assert - assert len(result['metric']) == 1 + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert len(result) == 1 assert result['metric'].values[0] == 'mae' assert result['model_id'].values[0] == 'test_model_id' assert result['timestamp'].values[0] == '2023-05-26 11:12:28' assert result['data_size'].values[0] == 2 assert result['interval_minutes'].values[0] == 5 - model_metrics_activity.info.assert_called_once_with( - "Calculating simple metrics for model test_model_id: ['mae']", metadata['metadata'] - ) + mock_instance.calculate.assert_called_once_with(['mae']) def test_calculate_simple_metrics_success_r2_only(model_metrics_activity): - # Arrange target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], @@ -847,24 +863,24 @@ def test_calculate_simple_metrics_success_r2_only(model_metrics_activity): 'interval_minutes': 5, } - # Act - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [{'metric': 'r2', 'value': 0.95}] + ) - # Assert - assert len(result['metric']) == 1 + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert len(result) == 1 assert result['metric'].values[0] == 'r2' assert result['model_id'].values[0] == 'test_model_id' assert result['timestamp'].values[0] == '2023-05-26 11:12:28' assert result['data_size'].values[0] == 2 assert result['interval_minutes'].values[0] == 5 - model_metrics_activity.info.assert_called_once_with( - "Calculating simple metrics for model test_model_id: ['r2']", metadata['metadata'] - ) + mock_instance.calculate.assert_called_once_with(['r2']) -def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity): - # Arrange - # All target values are the same, so ss_tot will be 0 +def test_calculate_simple_metrics_r2_zero_variance_delegates_to_lib(model_metrics_activity): + """r2 zero-variance is now the lib's responsibility. Activity just passes through.""" target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], @@ -881,24 +897,18 @@ def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity): 'interval_minutes': 5, } - # Act - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) - - # Assert - assert len(result['metric']) == 1 - assert result['metric'].values[0] == 'r2' - assert result['value'].values[0] == 0.0 # Should return 0.0 when ss_tot == 0 - assert result['model_id'].values[0] == 'test_model_id' - assert result['timestamp'].values[0] == '2023-05-26 11:12:28' - assert result['data_size'].values[0] == 2 - assert result['interval_minutes'].values[0] == 5 - model_metrics_activity.info.assert_called_once_with( - "Calculating simple metrics for model test_model_id: ['r2']", metadata['metadata'] + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [{'metric': 'r2', 'value': 0.0}] ) + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert result['value'].values[0] == 0.0 + mock_instance.calculate.assert_called_once_with(['r2']) + def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity): - # Arrange target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'], @@ -915,23 +925,26 @@ def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_ 'interval_minutes': 5, } - # Act - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) - - # Assert - assert len(result['metric']) == 2 - assert 'rmse' in result['metric'].values - assert 'mae' in result['metric'].values - assert all(model_id == 'test_model_id' for model_id in result['model_id'].values) - assert all(timestamp == '2023-05-26 11:12:29' for timestamp in result['timestamp'].values) - assert all(data_size == 3 for data_size in result['data_size'].values) - assert all(interval_minutes == 5 for interval_minutes in result['interval_minutes'].values) - model_metrics_activity.info.assert_called_once_with( - "Calculating simple metrics for model test_model_id: ['rmse', 'mae']", metadata['metadata'] + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [ + {'metric': 'rmse', 'value': 0.1}, + {'metric': 'mae', 'value': 0.1}, + ] ) + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) -def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity): + assert len(result) == 2 + assert set(result['metric'].values) == {'rmse', 'mae'} + assert all(mid == 'test_model_id' for mid in result['model_id'].values) + assert all(ts == '2023-05-26 11:12:29' for ts in result['timestamp'].values) + assert all(ds == 3 for ds in result['data_size'].values) + assert all(im == 5 for im in result['interval_minutes'].values) + mock_instance.calculate.assert_called_once_with(['rmse', 'mae']) + + +def test_calculate_simple_metrics_unknown_metric_raises(model_metrics_activity): target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], @@ -948,7 +961,79 @@ def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity) 'interval_minutes': 5, } - result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + mock_instance = MagicMock() + mock_instance.calculate.side_effect = ValueError('Unknown metric: unknown_metric') + mock_class = MagicMock(return_value=mock_instance) + mock_class.is_r2_supported = MagicMock(return_value=True) - assert len(result['metric']) == 1 + with patch('laborious.activities.model_metrics.RegressionMetrics', mock_class): + with raises(ValueError, match='Unknown metric'): + model_metrics_activity.calculate_simple_metrics(input_data) + + +def test_calculate_simple_metrics_r2_skipped_for_nonlinear_model(model_metrics_activity): + """When model_type is non-linear, r2 is excluded and a warning notification fires.""" + target_data = DataFrame( + { + 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], + 'target': [1.0, 2.0], + 'prediction': [1.1, 2.1], + } + ) + + input_data = { + **metadata, + 'model_id': 'test_model_id', + 'target_data': target_data.to_dict(), + 'metrics': ['rmse', 'r2'], + 'interval_minutes': 5, + 'model_type': 'XGBoost', + } + + mock_instance = MagicMock() + mock_instance.calculate.return_value = [{'metric': 'rmse', 'value': 0.1}] + mock_class = MagicMock(return_value=mock_instance) + mock_class.is_r2_supported = MagicMock(return_value=False) + + with patch('laborious.activities.model_metrics.RegressionMetrics', mock_class): + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert len(result) == 1 assert result['metric'].values[0] == 'rmse' + mock_class.is_r2_supported.assert_called_once_with('XGBoost') + mock_instance.calculate.assert_called_once_with(['rmse']) + model_metrics_activity.warning.assert_called_once() + model_metrics_activity.send_notification.assert_called_once() + call_kwargs = model_metrics_activity.send_notification.call_args.kwargs + assert call_kwargs['notification_id'] == 'SIMPLE_METRICS_R2_UNSUPPORTED' + assert call_kwargs['level'] == NotificationLevel.WARNING + + +def test_calculate_simple_metrics_no_model_type_includes_r2(model_metrics_activity): + """When model_type is None (legacy input), r2 is included without check.""" + target_data = DataFrame( + { + 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], + 'target': [1.0, 2.0], + 'prediction': [1.1, 2.1], + } + ) + + input_data = { + **metadata, + 'model_id': 'test_model_id', + 'target_data': target_data.to_dict(), + 'metrics': ['r2'], + 'interval_minutes': 5, + # no model_type key + } + + patcher, mock_class, mock_instance = _mock_regression_metrics_class( + [{'metric': 'r2', 'value': 0.95}] + ) + + with patcher: + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) + + assert result['metric'].values[0] == 'r2' + mock_class.is_r2_supported.assert_not_called()