Skip to content

Commit 166b751

Browse files
committed
turn expressions into tendencies too
Replace DerivedWaveform with an ExpressionTendency on the single Waveform class, integrated with develop's import/static waveforms. Bare scalars are now classified by content: a number is a constant, a string that references another waveform or uses an operator/function is an expression, and a plain word is a literal string constant. Explicit {value:}/{expression:} override the heuristic. Example config uses the compact bare form.
1 parent 2ff4794 commit 166b751

15 files changed

Lines changed: 417 additions & 366 deletions

tests/test_configuration.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -255,6 +255,7 @@ def test_dump_comments():
255255

256256
yaml_str = dedent("""
257257
globals:
258+
version: 2
258259
dd_version: 3.42.0
259260
imports:
260261
ec_launchers: imas:hdf5?path=test_md
@@ -284,6 +285,7 @@ def test_dump_globals():
284285
dumped_yaml = config.dump()
285286
expected_dump = dedent("""
286287
globals:
288+
version: 2
287289
dd_version: 3.41.0
288290
imports:
289291
ec_launchers: imas:mdsplus?path=test

tests/test_derived_waveform.py

Lines changed: 0 additions & 153 deletions
This file was deleted.

tests/test_expression.py

Lines changed: 109 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,109 @@
1+
import numpy as np
2+
import pytest
3+
4+
from waveform_editor.configuration import WaveformConfiguration
5+
from waveform_editor.waveform import Waveform
6+
7+
8+
@pytest.fixture
9+
def config():
10+
config = WaveformConfiguration()
11+
config.add_group("root_group", [])
12+
return config
13+
14+
15+
@pytest.fixture
16+
def filled_config(config):
17+
waveform_list = [
18+
{
19+
"user_type": "linear",
20+
"user_from": 10,
21+
"user_to": 20,
22+
"user_start": 5,
23+
"user_end": 15,
24+
"line_number": 1,
25+
}
26+
]
27+
waveform = Waveform(waveform=waveform_list, name="waveform/1")
28+
config.add_waveform(waveform, ["root_group"])
29+
return config
30+
31+
32+
def make_expression(config, name, expr, add=True):
33+
"""Create an expression waveform bound to ``config``."""
34+
waveform = Waveform(
35+
waveform=[{"user_expression": expr, "line_number": 0}],
36+
name=name,
37+
config=config,
38+
)
39+
if add:
40+
config.add_waveform(waveform, ["root_group"])
41+
return waveform
42+
43+
44+
def test_constant_expression(config):
45+
waveform = make_expression(config, "waveform/1", "3")
46+
assert waveform.is_expression
47+
assert not waveform.is_categorical
48+
assert waveform.dependencies == set()
49+
_, value = waveform.get_value(np.linspace(0, 100, 101))
50+
assert np.all(value == 3)
51+
52+
53+
def test_dependent_waveform(filled_config):
54+
waveform = make_expression(filled_config, "waveform/2", '"waveform/1"')
55+
assert waveform.dependencies == {"waveform/1"}
56+
time, value = waveform.get_value()
57+
assert time[0] == 5
58+
assert time[-1] == 15
59+
assert value[0] == 10
60+
assert value[-1] == 20
61+
_, value = waveform.get_value(np.array([0, 5, 10, 15, 20]))
62+
assert np.all(value == [10, 10, 15, 20, 20])
63+
64+
65+
def test_dependent_waveform_calc(filled_config):
66+
waveform = make_expression(filled_config, "waveform/2", '"waveform/1" * 10')
67+
assert waveform.dependencies == {"waveform/1"}
68+
_, value = waveform.get_value(np.array([0, 5, 10, 15, 20]))
69+
assert np.all(value == [100, 100, 150, 200, 200])
70+
71+
72+
def test_dependent_waveform_numpy(filled_config):
73+
waveform = make_expression(
74+
filled_config, "waveform/2", 'maximum("waveform/1" * 10, 150)'
75+
)
76+
assert waveform.dependencies == {"waveform/1"}
77+
_, value = waveform.get_value(np.array([0, 5, 10, 15, 20]))
78+
assert np.all(value == [150, 150, 150, 200, 200])
79+
80+
81+
def test_rename_dependency(filled_config):
82+
waveform = make_expression(filled_config, "waveform/2", '"waveform/1"', add=False)
83+
assert waveform.dependencies == {"waveform/1"}
84+
waveform.rename_dependency("waveform/1", "waveform/3")
85+
assert waveform.dependencies == {"waveform/3"}
86+
87+
88+
def test_function_access_control(filled_config):
89+
test_exprs = [
90+
('max("waveform/1")', False),
91+
('sum("waveform/1")', False),
92+
('eval("waveform/1")', False),
93+
('dot("waveform/1", "waveform/1")', False),
94+
('linalg.norm("waveform/1")', False),
95+
('linalg.inv("waveform/1")', False),
96+
('sin("waveform/1")', True),
97+
('log("waveform/1" + 1)', True),
98+
('maximum("waveform/1", 10)', True),
99+
]
100+
101+
time = np.linspace(filled_config.start, filled_config.end, 100)
102+
for expr, allowed in test_exprs:
103+
waveform = make_expression(filled_config, "waveform/2", expr, add=False)
104+
if allowed:
105+
_, result = waveform.get_value(time)
106+
assert result is not None
107+
else:
108+
with pytest.raises(NameError):
109+
waveform.get_value(time)

tests/test_yaml/example.yaml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
globals:
2+
version: 2
23
dd_version: 4.0.0
34
imports: {}
45
dummy_waveform:

tests/test_yaml_parser.py

Lines changed: 35 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
from pytest import approx
33

44
from waveform_editor.configuration import WaveformConfiguration
5-
from waveform_editor.derived_waveform import DerivedWaveform
5+
from waveform_editor.static_waveform import StaticWaveform
66
from waveform_editor.tendencies.constant import ConstantTendency
77
from waveform_editor.tendencies.linear import LinearTendency
88
from waveform_editor.tendencies.periodic.sawtooth_wave import SawtoothWaveTendency
@@ -132,19 +132,47 @@ def test_scientific_notation(yaml_parser):
132132

133133

134134
def test_constant_shorthand_notation(yaml_parser):
135-
"""Test if shorthand notation is parsed correctly."""
135+
"""A bare scalar is shorthand for a single constant tendency."""
136136

137137
waveforms = {"waveform: 5": 5, "waveform: 1.23": 1.23}
138138

139-
for waveform, expected_value in waveforms.items():
140-
waveform = yaml_parser.parse_waveform(waveform)
141-
assert isinstance(waveform, DerivedWaveform)
142-
assert waveform.yaml == expected_value
143-
assert not waveform.annotations
139+
for waveform_str, expected_value in waveforms.items():
140+
waveform = yaml_parser.parse_waveform(waveform_str)
141+
assert isinstance(waveform, Waveform)
142+
assert not waveform.is_expression
144143
assert waveform.dependencies == set()
144+
assert waveform.tendencies[0].value == expected_value
145+
assert not waveform.annotations
145146
assert not yaml_parser.parse_errors
146147

147148

149+
def test_bare_word_is_literal_string(yaml_parser):
150+
"""A bare plain word is a literal string constant, not an expression."""
151+
waveform = yaml_parser.parse_waveform("waveform: total")
152+
assert isinstance(waveform, StaticWaveform)
153+
assert not waveform.is_expression
154+
assert waveform.value == "total"
155+
assert not yaml_parser.parse_errors
156+
157+
158+
def test_bare_number_is_constant(yaml_parser):
159+
"""A bare number is a constant tendency, not an expression."""
160+
waveform = yaml_parser.parse_waveform("waveform: 3.5")
161+
assert isinstance(waveform, Waveform)
162+
assert not waveform.is_expression
163+
assert waveform.tendencies[0].value == 3.5
164+
assert not yaml_parser.parse_errors
165+
166+
167+
def test_bare_reference_is_expression(yaml_parser):
168+
"""A bare string that references another waveform (quoted) is an expression."""
169+
waveform = yaml_parser.parse_waveform('waveform: \'"other/1" * 2\'')
170+
assert isinstance(waveform, Waveform)
171+
assert waveform.is_expression
172+
assert waveform.dependencies == {"other/1"}
173+
assert not yaml_parser.parse_errors
174+
175+
148176
def test_load_yaml(config):
149177
"""Test if yaml is loaded correctly."""
150178
yaml_str = """

waveform_editor/base_waveform.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,23 @@ def get_value(
2828
def get_yaml_string(self) -> str:
2929
raise NotImplementedError
3030

31+
@property
32+
def dependencies(self):
33+
"""Names of other waveforms this waveform depends on. Empty unless the waveform
34+
contains expressions (overridden by :class:`Waveform`)."""
35+
return set()
36+
37+
@property
38+
def is_expression(self):
39+
"""Whether this waveform is computed from an expression. False by default."""
40+
return False
41+
42+
def prepare_expression(self): # noqa: B027
43+
"""Re-parse expression tendencies. No-op for waveforms without expressions."""
44+
45+
def rename_dependency(self, old_name, new_name): # noqa: B027
46+
"""Rename a referenced waveform. No-op for waveforms without expressions."""
47+
3148
def get_metadata(self, dd_version):
3249
"""Parses the name of the waveform and returns the IDS metadata for this
3350
waveform. The name must be formatted as follows: ``<IDS-Name>/<IDS-path>``

0 commit comments

Comments
 (0)