SIENTIAPDE-1255: Add tests to ensure get_required_columns removes columns from self_operations and cross_operations.

This commit is contained in:
Bruno Domingues
2025-10-20 22:27:58 -03:00
parent 614eb0cb7f
commit 3e52253a98

View File

@@ -6,7 +6,23 @@ import numpy as np
import pandas as pd import pandas as pd
from pytest import raises 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 # 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 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(): 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.""" """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}) preprocessor = DataPreprocessor(created_lags={'var1': 1, 'var2': 1})