"""Unit tests for TrainModelResult dataclass.""" from unittest.mock import MagicMock import pandas as pd import pytest from model_manager.utils.models.train_model_params import TrainModelParams from model_manager.utils.models.train_model_result import TrainModelResult @pytest.fixture def sample_params(): """Create sample TrainModelParams for testing.""" return TrainModelParams( variable_columns=['var1', 'var2'], lag_train={'var1': 5, 'var2': 5}, lag_val={'var1': 3, 'var2': 3}, target_variable='target', rem_static_win=True, low_lim={'var1': 0.0, 'var2': 0.0}, upp_lim={'var1': 100.0, 'var2': 100.0}, window=10, use_scaler=True, include_ar=False, bucket_name='test-bucket', file_name='test-file.csv', line_separator='\n', decimal_separator='.', train_size=80, shuffle=True, experiment_run_id=123, experiment_name='test-experiment', removed_intervals=[], model_name='Linear Regression', degree=1, interaction_only=False, nan_treatment='drop', start_date=None, end_date=None, scaler_name='Standard Scaler', support_filters={}, ) @pytest.fixture def sample_dataframes(): """Create sample DataFrames for testing.""" X_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [4, 5, 6]}) X_test = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10]}) y_train = pd.DataFrame({'target': [10, 20, 30]}) y_test = pd.DataFrame({'target': [40, 50]}) return X_train, X_test, y_train, y_test def test_train_model_result_creation(sample_params, sample_dataframes): """Test creating TrainModelResult with required fields.""" x_train, x_test, y_train, y_test = sample_dataframes process_data = MagicMock() regr = MagicMock() scaler_dict = {'var1': {'min': 0, 'max': 100}} result = TrainModelResult( params=sample_params, process_data=process_data, x_train=x_train, x_test=x_test, y_train=y_train, y_test=y_test, regr=regr, scaler_dict=scaler_dict, ) assert result.params == sample_params assert result.process_data == process_data assert result.x_train.equals(x_train) assert result.x_test.equals(x_test) assert result.y_train.equals(y_train) assert result.y_test.equals(y_test) assert result.regr == regr assert result.scaler_dict == scaler_dict def test_train_model_result_optional_fields_default_none(sample_params, sample_dataframes): """Test that optional fields default to None.""" x_train, x_test, y_train, y_test = sample_dataframes result = TrainModelResult( params=sample_params, process_data=MagicMock(), x_train=x_train, x_test=x_test, y_train=y_train, y_test=y_test, regr=MagicMock(), scaler_dict={}, ) assert result.y_pred is None assert result.mse_val is None assert result.mae_val is None assert result.r2_val is None assert result.run_name is None assert result.report_path is None assert result.train_data_path is None assert result.test_data_path is None assert result.run_dir is None def test_train_model_result_with_metrics(sample_params, sample_dataframes): """Test TrainModelResult with metrics populated.""" x_train, x_test, y_train, y_test = sample_dataframes y_pred = pd.Series([41, 49]) result = TrainModelResult( params=sample_params, process_data=MagicMock(), x_train=x_train, x_test=x_test, y_train=y_train, y_test=y_test, regr=MagicMock(), scaler_dict={}, y_pred=y_pred, mse_val=1.5, mae_val=1.2, r2_val=0.95, ) assert result.y_pred.equals(y_pred) assert result.mse_val == 1.5 assert result.mae_val == 1.2 assert result.r2_val == 0.95 def test_train_model_result_with_artifact_paths(sample_params, sample_dataframes): """Test TrainModelResult with artifact paths populated.""" x_train, x_test, y_train, y_test = sample_dataframes result = TrainModelResult( params=sample_params, process_data=MagicMock(), x_train=x_train, x_test=x_test, y_train=y_train, y_test=y_test, regr=MagicMock(), scaler_dict={}, run_name='test-experiment-1', report_path='/path/to/report.html', train_data_path='/path/to/train_data.csv', test_data_path='/path/to/test_data.csv', run_dir='/path/to/run_dir', ) assert result.run_name == 'test-experiment-1' assert result.report_path == '/path/to/report.html' assert result.train_data_path == '/path/to/train_data.csv' assert result.test_data_path == '/path/to/test_data.csv' assert result.run_dir == '/path/to/run_dir' def test_train_model_result_is_dataclass(sample_params, sample_dataframes): """Test that TrainModelResult is a dataclass.""" x_train, x_test, y_train, y_test = sample_dataframes result = TrainModelResult( params=sample_params, process_data=MagicMock(), x_train=x_train, x_test=x_test, y_train=y_train, y_test=y_test, regr=MagicMock(), scaler_dict={}, ) # Dataclasses have __dataclass_fields__ attribute assert hasattr(result, '__dataclass_fields__') assert 'params' in result.__dataclass_fields__ assert 'process_data' in result.__dataclass_fields__ assert 'x_train' in result.__dataclass_fields__ def test_train_model_result_field_count(): """Test that TrainModelResult has exactly 20 fields.""" from dataclasses import fields result_fields = fields(TrainModelResult) assert len(result_fields) == 20 field_names = {f.name for f in result_fields} expected_fields = { 'params', 'process_data', 'x_train', 'x_test', 'y_train', 'y_test', 'regr', 'scaler_dict', 'y_pred', 'y_train_pred', 'mse_val', 'mae_val', 'r2_val', 'equation', 'equation_path', 'run_name', 'report_path', 'train_data_path', 'test_data_path', 'run_dir', } assert field_names == expected_fields def test_train_model_result_complete_workflow(sample_params, sample_dataframes): """Test TrainModelResult through a complete workflow simulation.""" x_train, x_test, y_train, y_test = sample_dataframes # Step 1: Create result after training result = TrainModelResult( params=sample_params, process_data=MagicMock(), x_train=x_train, x_test=x_test, y_train=y_train, y_test=y_test, regr=MagicMock(), scaler_dict={'var1': {'min': 0, 'max': 100}}, ) # Step 2: Add predictions and metrics result.y_pred = pd.Series([41, 49]) result.mse_val = 1.5 result.mae_val = 1.2 result.r2_val = 0.95 # Step 3: Add artifact paths result.run_name = 'test-experiment-1' result.report_path = '/path/to/report.html' result.train_data_path = '/path/to/train_data.csv' result.test_data_path = '/path/to/test_data.csv' result.run_dir = '/path/to/run_dir' # Verify all fields are populated assert result.y_pred is not None assert result.mse_val == 1.5 assert result.mae_val == 1.2 assert result.r2_val == 0.95 assert result.run_name == 'test-experiment-1' assert result.report_path == '/path/to/report.html' assert result.train_data_path == '/path/to/train_data.csv' assert result.test_data_path == '/path/to/test_data.csv' assert result.run_dir == '/path/to/run_dir'