|
| 1 | +from datetime import datetime |
| 2 | + |
| 3 | +import numpy as np |
| 4 | + |
| 5 | +from libecalc.common.time_utils import Period |
| 6 | +from libecalc.domain.regularity import Regularity |
| 7 | +from libecalc.presentation.yaml.domain.time_series_mask import TimeSeriesMask |
| 8 | +from libecalc.presentation.yaml.domain.time_series_expression import TimeSeriesExpression |
| 9 | +from libecalc.presentation.yaml.domain.expression_time_series_variable import ExpressionTimeSeriesVariable |
| 10 | + |
| 11 | + |
| 12 | +periods = [ |
| 13 | + Period(start=datetime(2020, 1, 1), end=datetime(2021, 1, 1)), |
| 14 | + Period(start=datetime(2021, 1, 1), end=datetime(2022, 1, 1)), |
| 15 | + Period(start=datetime(2022, 1, 1), end=datetime(2023, 1, 1)), |
| 16 | +] |
| 17 | + |
| 18 | + |
| 19 | +def test_time_series_mask_none(): |
| 20 | + """Test that TimeSeriesMask with None mask returns the input array unchanged.""" |
| 21 | + mask = TimeSeriesMask.from_evaluated_mask(None) |
| 22 | + arr = np.array([1, 2, 3]) |
| 23 | + np.testing.assert_array_equal(mask.apply(arr), arr) |
| 24 | + |
| 25 | + |
| 26 | +def test_time_series_mask_int_mask(): |
| 27 | + """Test that TimeSeriesMask with integer mask applies masking correctly.""" |
| 28 | + mask = TimeSeriesMask.from_evaluated_mask(np.array([1, 0, 1])) |
| 29 | + arr = np.array([10, 20, 30]) |
| 30 | + np.testing.assert_array_equal(mask.apply(arr), [10, 0, 30]) |
| 31 | + |
| 32 | + |
| 33 | +def test_time_series_mask_float_mask(): |
| 34 | + """Test that TimeSeriesMask with float mask applies masking (nonzero as 1, zero as 0).""" |
| 35 | + mask = TimeSeriesMask.from_evaluated_mask(np.array([0.5, 0.0, -2.1])) |
| 36 | + arr = np.array([7, 8, 9]) |
| 37 | + np.testing.assert_array_equal(mask.apply(arr), [7, 0, 9]) |
| 38 | + |
| 39 | + |
| 40 | +def test_time_series_mask_condition_mask(expression_evaluator_factory): |
| 41 | + """Test that TimeSeriesExpression uses condition mask correctly.""" |
| 42 | + evaluator = expression_evaluator_factory.from_periods( |
| 43 | + periods, variables={"SIM1;GAS_PROD": [10.0, 5.0, 10.0], "TEST_VAR": [1.0, 2.0, 3.0]} |
| 44 | + ) |
| 45 | + |
| 46 | + expr = TimeSeriesExpression(expressions="TEST_VAR", expression_evaluator=evaluator, condition="SIM1;GAS_PROD > 5") |
| 47 | + # The condition "SIM1;GAS_PROD > 5" results in a mask [1, 0, 1], which is applied to the values of TEST_VAR |
| 48 | + assert expr.get_masked_values() == [1, 0, 3] |
| 49 | + |
| 50 | + |
| 51 | +def test_expression_time_series_variable_get_values_with_mask(expression_evaluator_factory): |
| 52 | + """Test that ExpressionTimeSeriesVariable applies mask to values as expected.""" |
| 53 | + evaluator = expression_evaluator_factory.from_periods( |
| 54 | + periods, variables={"SIM1;GAS_PROD": [10.0, 5.0, 10.0], "TEST_VAR": [1.0, 2.0, 3.0]} |
| 55 | + ) |
| 56 | + regularity = Regularity(expression_evaluator=evaluator, target_period=evaluator.get_period(), expression_input=1) |
| 57 | + expr = TimeSeriesExpression(expressions="TEST_VAR", expression_evaluator=evaluator, condition="SIM1;GAS_PROD > 5") |
| 58 | + |
| 59 | + var = ExpressionTimeSeriesVariable(name="test", time_series_expression=expr, regularity=regularity, is_rate=False) |
| 60 | + # The mask [1, 0, 1], created from condition, is applied to [1.0, 2.0, 3.0], resulting in [1.0, 0.0, 3.0] |
| 61 | + assert var.get_values() == [1, 0, 3] |
0 commit comments