From 3e52253a980148056f4b87914302dc041673a60d Mon Sep 17 00:00:00 2001 From: Bruno Domingues Date: Mon, 20 Oct 2025 22:27:58 -0300 Subject: [PATCH] SIENTIAPDE-1255: Add tests to ensure get_required_columns removes columns from self_operations and cross_operations. --- tests/sientia/test_models.py | 43 +++++++++++++++++++++++++++++++++++- 1 file changed, 42 insertions(+), 1 deletion(-) diff --git a/tests/sientia/test_models.py b/tests/sientia/test_models.py index d84fddd..2752eec 100644 --- a/tests/sientia/test_models.py +++ b/tests/sientia/test_models.py @@ -6,7 +6,23 @@ import numpy as np import pandas as pd from pytest import raises -from model_manager.sientia.models import DataPreprocessor, LinearRegressionModel +from model_manager.sientia.models import ( + DataPreprocessor, + LinearRegressionModel, +) + + +class _IterableWithContains: + def __init__(self, iterable, contains_values): + self._iterable = iterable + self._contains = set(contains_values) + + def __iter__(self): + return iter(self._iterable) + + def __contains__(self, item): + return item in self._contains + # LinearRegressionModel Tests @@ -573,6 +589,31 @@ def test_data_preprocessor_get_required_columns_removes_duplicates(): assert 'var1' not in result or result.count('var1') <= 1 +def test_data_preprocessor_get_required_columns_removes_self_operations_branch(): + """Ensure line 278 removes columns present in self_operations iterable.""" + preprocessor = DataPreprocessor( + self_operations=_IterableWithContains(['{var1}_{pow}_{2}'], contains_values=['var1']) + ) + existing_columns: list[str] = [] + + result = preprocessor.get_required_columns(existing_columns) + + assert 'var1' not in result + + +def test_data_preprocessor_get_required_columns_removes_cross_operations_branch(): + """Ensure line 285 removes columns present in cross_operations iterable.""" + preprocessor = DataPreprocessor( + cross_operations=_IterableWithContains(['{var1}_{*}_{var2}'], contains_values=['var1']) + ) + existing_columns: list[str] = [] + + result = preprocessor.get_required_columns(existing_columns) + + assert 'var1' not in result + assert 'var2' in result + + def test_data_preprocessor_get_required_columns_removes_created_lags(): """Test get_required_columns removes columns from created_lags when column is in created_lags dict - covers line 292.""" preprocessor = DataPreprocessor(created_lags={'var1': 1, 'var2': 1})