Skip to content

Commit 695ce66

Browse files
committed
test: add unit tests
1 parent b7eafae commit 695ce66

3 files changed

Lines changed: 71 additions & 4 deletions

File tree

src/libecalc/presentation/yaml/domain/expression_time_series_variable.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,8 +25,6 @@ def __init__(
2525
self._regularity = regularity
2626
self._is_rate = is_rate
2727

28-
self._condition = self._time_series_expression.get_condition_mask()
29-
3028
@property
3129
def name(self) -> str:
3230
return self._name
@@ -36,7 +34,7 @@ def is_rate(self) -> bool:
3634
return self._is_rate
3735

3836
def get_values(self) -> list[float]:
39-
values: np.ndarray = np.asarray(self._time_series_expression.get_evaluated_expressions(), dtype=np.float64)
37+
values: np.ndarray = np.asarray(self._time_series_expression.get_masked_values(), dtype=np.float64)
4038
# If some of these are rates, we need to calculate stream day rate for use
4139
# Also take a copy of the calendar day rate and stream day rate for input to result object
4240

@@ -46,7 +44,6 @@ def get_values(self) -> list[float]:
4644
regularity=self._regularity.values,
4745
)
4846

49-
values = self._condition.apply(values)
5047
return values.tolist()
5148

5249
def get_periods(self) -> Periods:

src/libecalc/presentation/yaml/domain/time_series_expression.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,15 @@ def get_evaluated_expressions(self) -> list[float]:
5454
arr = np.atleast_1d(arr[0]) # Flattens (1, N) to (N,)
5555
return arr.tolist()
5656

57+
def get_masked_values(self) -> list[float]:
58+
"""
59+
Returns the evaluated expressions with the condition mask applied.
60+
"""
61+
values = np.asarray(self.get_evaluated_expressions(), dtype=np.float64)
62+
mask = self.get_condition_mask()
63+
masked = mask.apply(values)
64+
return masked.tolist()
65+
5766
def get_condition_mask(self) -> TimeSeriesMask:
5867
mask = self.expression_evaluator.evaluate(expression=self._condition) if self._condition is not None else None
5968
return TimeSeriesMask.from_evaluated_mask(mask)
Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,61 @@
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

Comments
 (0)