diff --git a/pyaml/arrays/element_array.py b/pyaml/arrays/element_array.py index 26b7ce26..428b8e1d 100644 --- a/pyaml/arrays/element_array.py +++ b/pyaml/arrays/element_array.py @@ -13,6 +13,7 @@ from ..bpm.bpm import BPM from ..common.element import Element, __pyaml_repr__ from ..common.exception import PyAMLException +from ..common.name_matching import resolve_names from ..magnet.cfm_magnet import CombinedFunctionMagnet from ..magnet.magnet import Magnet from ..magnet.serialized_magnet import SerializedMagnets @@ -240,22 +241,6 @@ def mro_as_list(cls: type) -> list[type]: return self.__create_array("", chosen, elements) - def _select_names(self, pattern: str) -> "ElementArray": - """Select names without interpreting field selectors. - - Parameters - ---------- - pattern : str - A fnmatch pattern applied to each element name. - - Returns - ------- - ElementArray - Typed selection in the original order, including an empty array - when no names match. - """ - return self._typed_array([element for element in self if fnmatch.fnmatch(element.get_name(), pattern)]) - def __is_bool_mask(self, other: object) -> bool: """Return True if 'other' looks like a boolean mask (list or numpy array).""" # --- numpy boolean array --- @@ -576,8 +561,15 @@ def __getitem__(self, key): Parameters ---------- - key : int, slice or str - Element index, slice, name pattern, or existing field selector. + key : int, slice, str, list[str] or tuple[str, ...] + Element index, slice, name pattern (or ``field:pattern`` field + selector), or a list/tuple of name patterns. A name pattern is a + literal name (must match an element in this array), an fnmatch + wildcard (``*``, ``?`` or ``[``), or a ``re:``-prefixed regular + expression. A list or tuple resolves each entry independently + and unions the results. Prefix a name pattern with ``~`` to + exclude its matches instead; a lone ``~pattern`` means every + element except those matches. Returns ------- @@ -585,33 +577,47 @@ def __getitem__(self, key): Indexed element or an array inferred from all selected elements. Different Magnet subclasses produce a MagnetArray; mixed element families produce an ElementArray. Empty selections return an - empty ElementArray. + empty ElementArray, except a literal name pattern (or literal + entry within a list/tuple) matching nothing, which raises. + + Raises + ------ + PyAMLException + If a literal name pattern, or a literal entry within a list or + tuple, matches no element in this array, or a ``re:`` pattern is + not a valid regular expression. Examples -------- >>> magnets = sr.design.magnets.get() >>> subset = magnets[:] # MagnetArray, including mixed magnet classes >>> correctors = magnets["SH*"] # MagnetArray + >>> correctors = magnets["re:^SH1A-C0[12]-H$"] # MagnetArray + >>> selection = magnets[["SH1A-C01-H", "SH*-V"]] # Union of patterns + >>> all_but_one = magnets["~SH1A-C01-H"] # Every magnet except one """ if isinstance(key, slice): # Slicing r = super().__getitem__(key) + elif isinstance(key, (list, tuple)): + # Selection by a list/tuple of name patterns + matched = set(resolve_names([e.get_name() for e in self], key)) + r = [e for e in self if e.get_name() in matched] + elif isinstance(key, str): - fields = key.split(":") + fields = [] if key.startswith("re:") else key.split(":") if len(fields) <= 1: - # Selection by name - r = [] - for e in self: - if fnmatch.fnmatch(e.get_name(), key): - r.append(e) + # Selection by name pattern + matched = set(resolve_names([e.get_name() for e in self], key)) + r = [e for e in self if e.get_name() in matched] else: # Selection by fields r = [] for e in self: txt = self.__eval_field(fields[0], e) - if fnmatch.fnmatch(txt, fields[1]): + if fnmatch.fnmatchcase(txt, fields[1]): r.append(e) else: diff --git a/pyaml/common/holders/diagnostic_holder.py b/pyaml/common/holders/diagnostic_holder.py index 361a06ed..588b74c7 100644 --- a/pyaml/common/holders/diagnostic_holder.py +++ b/pyaml/common/holders/diagnostic_holder.py @@ -6,6 +6,7 @@ from ...diagnostics.tune_monitor import BetatronTuneMonitor from ..element import Element, __pyaml_repr__ from ..exception import PyAMLException +from ..name_matching import is_wildcard, resolve_names from .sub_holders import BPMHolder, BPMsHolder if TYPE_CHECKING: @@ -95,6 +96,37 @@ def get(self, name: str = None) -> "Element | ElementArray": return ElementArray("", list(self._peer._DIAG.values())) return self._peer._get_diagnostic(name) + def __getitem__(self, key: str | list[str] | tuple[str, ...]) -> "Element | ElementArray": + """ + Return a diagnostic, or a selection typed as a generic ElementArray. + + Parameters + ---------- + key : str, list[str] or tuple[str, ...] + An exact literal name returns the stored diagnostic. An fnmatch + wildcard, a ``re:``-prefixed regular expression, or a list/tuple + of such patterns returns an ElementArray of matches (possibly + empty). + + Returns + ------- + Element or ElementArray + The stored diagnostic for an exact literal name, otherwise an + ElementArray of matches. + + Raises + ------ + PyAMLException + If an exact literal name, or a literal entry within a list or + tuple, does not match any diagnostic, or a ``re:`` pattern is + not a valid regular expression. + """ + store = self._peer._DIAG + if isinstance(key, str) and not key.startswith(("re:", "~")) and not is_wildcard(key): + return self._peer._get_diagnostic(key) + names = resolve_names(store.keys(), key, what="Diagnostic") + return ElementArray("", [store[n] for n in names]) + @property def betatron_tune(self) -> BetatronTuneMonitor: """ diff --git a/pyaml/common/holders/element_holder.py b/pyaml/common/holders/element_holder.py index 2a0fb27e..78e1e71f 100644 --- a/pyaml/common/holders/element_holder.py +++ b/pyaml/common/holders/element_holder.py @@ -1,7 +1,5 @@ """Store, resolve, and group elements shared by runtime backends.""" -import fnmatch -import re from abc import ABCMeta, abstractmethod from typing import TYPE_CHECKING, overload @@ -17,6 +15,7 @@ from ..abstract_aggregator import ScalarAggregator from ..element import Element from ..exception import PyAMLException +from ..name_matching import is_wildcard, resolve_names from .diagnostic_holder import DiagnosticHolder from .rf_holder import RFHolder from .sub_holders import ( @@ -311,29 +310,33 @@ def create_bpm_aggregators(self, bpms: list[BPM]) -> list[ScalarAggregator | Non # Elements - def find_elements(self, filter: str) -> list[str]: + def find_elements(self, filter: str | list[str] | tuple[str, ...]) -> list[str]: """ - Find element names matching a literal, wildcard, or regular expression. + Find element names matching one or several literal, wildcard, or regex patterns. Parameters ---------- - filter : str - Pattern to match. Prefix with ``re:`` for a regular expression. + filter : str, list[str] or tuple[str, ...] + Pattern, or patterns, to match. A pattern is a literal name, an + fnmatch wildcard (``*``, ``?`` or ``[`` anywhere in the string), + or a regular expression prefixed with ``re:``. Prefix any of + those with ``~`` to exclude its matches instead, as in the + ``elements:`` selector list of a YAML array declaration, e.g. + ``["QD2*", "QF1*", "~QF1E-C05"]``. A lone ``~pattern`` (or a list + made only of ``~`` entries) means "every element except those". Returns ------- list[str] Matching element names. - """ - if filter.startswith("re:"): - pattern = re.compile(rf"{filter[3:]}") - elements = [k for k in self._ALL.keys() if pattern.fullmatch(k)] - elif "*" in filter or "?" in filter: - elements = [k for k in self._ALL.keys() if fnmatch.fnmatch(k, filter)] - else: - elements = [filter] - return elements + Raises + ------ + PyAMLException + If a literal pattern (or a ``~``-prefixed literal) matches no + element, or a ``re:`` pattern is not a valid regular expression. + """ + return resolve_names(self._ALL.keys(), filter, what="Element") def _fill_array( self, @@ -450,32 +453,46 @@ def __getitem__(self, key: int) -> Element: ... def __getitem__(self, key: slice) -> ElementArray: ... @overload - def __getitem__(self, key: str) -> Element | ElementArray | None: ... + def __getitem__(self, key: str) -> Element: ... + + @overload + def __getitem__(self, key: list[str] | tuple[str, ...]) -> ElementArray: ... - def __getitem__(self, key: int | slice | str) -> Element | ElementArray | None: + def __getitem__(self, key: int | slice | str | list[str] | tuple[str, ...]) -> Element | ElementArray: """Retrieve an element or select a collection. Parameters ---------- - key : int, slice or str - Index in registration order, slice, exact name, or name pattern. - Strings containing ``*``, ``?`` or ``[`` use fnmatch matching. - Other strings are exact registry keys. Colons are literal. + key : int, slice, str, list[str] or tuple[str, ...] + Index in registration order, slice, exact name, name pattern, or + a list/tuple of patterns. Strings containing ``*``, ``?`` or + ``[`` use fnmatch matching; a ``re:`` prefix uses a regular + expression instead. Any other string is an exact registry key. + Colons are literal. A list or tuple resolves each entry + independently and unions the results. Prefix any pattern with + ``~`` to exclude its matches instead, as in a YAML array's + ``elements:`` list; a lone ``~pattern`` means every element + except those matches. Returns ------- - Element or ElementArray or None - An index returns an element. An exact name returns its element - or None. Patterns and slices return the most specific compatible - array, or an empty ElementArray when nothing matches. + Element or ElementArray + An index or an exact literal name returns its element. Patterns, + lists/tuples of patterns, and slices return the most specific + compatible array, or an empty ElementArray when nothing matches. The full slice ``[:]`` returns a generic ElementArray, like get(). Raises ------ + PyAMLException + If an exact literal name, or a literal entry within a list or + tuple, does not match any registered element, or if a ``re:`` + pattern is not a valid regular expression. IndexError If the index is out of bounds. TypeError - If the key is neither an integer, a slice, nor a string. + If the key is neither an integer, a slice, a string, nor a + list/tuple of strings. ValueError If a slice has a zero step. @@ -483,25 +500,35 @@ def __getitem__(self, key: int | slice | str) -> Element | ElementArray | None: ----- Indices follow insertion order, not necessarily lattice order. Collections share element references but do not modify the registry. - Field filters and regular expressions are not interpreted here. + Field filters are not interpreted here. Examples -------- >>> bpm = sr.live["BPM01"] - >>> missing = sr.live["UNKNOWN"] # None + >>> missing = sr.live["UNKNOWN"] # raises PyAMLException >>> bpms = sr.live["BPM*"] >>> bpms = sr.live["BPM0[123]"] # BPM01, BPM02 or BPM03 >>> bpms = sr.live["BPM0[1-3]"] # Same selection using a range >>> quads = sr.live["Q[FD]*"] # Names starting with QF or QD >>> bpms = sr.live["BPM0[!3]"] # One character after BPM0, except 3 + >>> bpms = sr.live["re:^BPM0[12]$"] # Regular expression + >>> mixed = sr.live[["BPM01", "QF1*"]] # Union of several patterns + >>> all_but_one = sr.live["~BPM01"] # Every element except BPM01 + >>> most_quads = sr.live[["QF1*", "~QF1A-C01"]] # QF1* minus one name >>> first = sr.live[0] >>> subset = sr.live[1:10] >>> all_elements = sr.live[:] """ if isinstance(key, str): - if any(marker in key for marker in "*?["): - return self.get()._select_names(key) - return self._ALL.get(key) + if key.startswith("re:") or key.startswith("~") or is_wildcard(key): + names = resolve_names(self._ALL.keys(), key) + return self.get()._typed_array([self._ALL[n] for n in names]) + if key not in self._ALL: + raise PyAMLException(f"Element {key} not defined") + return self._ALL[key] + if isinstance(key, (list, tuple)): + names = resolve_names(self._ALL.keys(), key) + return self.get()._typed_array([self._ALL[n] for n in names]) if isinstance(key, int): return list(self._ALL.values())[key] if isinstance(key, slice): @@ -509,7 +536,7 @@ def __getitem__(self, key: int | slice | str) -> Element | ElementArray | None: if key == slice(None): return elements return elements._typed_array(list(elements)[key]) - raise TypeError("ElementHolder keys must be integers, slices or strings") + raise TypeError("ElementHolder keys must be integers, slices, strings, or lists/tuples of strings") def fill_element_array(self, arrayName: str, elementNames: list[str]): """ diff --git a/pyaml/common/holders/generic_array_holder.py b/pyaml/common/holders/generic_array_holder.py index 4004a52a..50fec2c1 100644 --- a/pyaml/common/holders/generic_array_holder.py +++ b/pyaml/common/holders/generic_array_holder.py @@ -110,17 +110,41 @@ def add(self, arrayName: str, elementNames: list[str]): def __getitem__(self, key): """ - Return an element from the aggregate array by index. + Select from the aggregate array of every individual element. + + Delegates to :meth:`ElementArray.__getitem__ + ` on ``self.get()`` + (the array of every individual element of this type), so ``key`` + matches against **individual element names**, not against the + registered array/family names that :meth:`get` searches. These are + deliberately two different, non-overlapping namespaces: ``get(name)`` + looks up a configured family (e.g. ``"QForTune"``), while ``[key]`` + looks up the elements themselves (e.g. ``"QF1A-C01"`` or ``"QF1*"``). Parameters ---------- - key : int or slice - Index or slice passed to the aggregate array. + key : int, slice, str, list[str] or tuple[str, ...] + Index or slice into the aggregate array, an individual element's + exact name, an fnmatch wildcard or ``re:`` regular expression + over element names, or a list/tuple of such patterns. Returns ------- object Element or sub-array selected by ``key``. + + Raises + ------ + PyAMLException + If ``key`` is an exact literal element name (or a literal entry + within a list or tuple) that matches no individual element, or a + ``re:`` pattern is not a valid regular expression. + + Examples + -------- + >>> family = sr.live.magnets.get("QForTune") # array-name namespace + >>> one_magnet = sr.live.magnets["QF1A-C01"] # element-name namespace + >>> some_magnets = sr.live.magnets["QF1*"] # element-name namespace """ return self.get().__getitem__(key) diff --git a/pyaml/common/holders/generic_element_holder.py b/pyaml/common/holders/generic_element_holder.py index a23dec8b..101796b9 100644 --- a/pyaml/common/holders/generic_element_holder.py +++ b/pyaml/common/holders/generic_element_holder.py @@ -2,7 +2,9 @@ from typing import TYPE_CHECKING, Generic, TypeVar +from ...arrays.element_array import ElementArray from ..element import Element, __pyaml_repr__ +from ..name_matching import is_wildcard, resolve_names if TYPE_CHECKING: from .element_holder import ElementHolder @@ -73,6 +75,37 @@ def get(self, name: str) -> T: """ return self._peer._get(self._what, name, self._store) + def __getitem__(self, key: str | list[str] | tuple[str, ...]) -> "T | ElementArray": + """ + Return an element, or a selection typed as a generic ElementArray. + + Parameters + ---------- + key : str, list[str] or tuple[str, ...] + An exact literal name returns the stored element. An fnmatch + wildcard (``*``, ``?`` or ``[``), a ``re:``-prefixed regular + expression, or a list/tuple of such patterns returns an + :class:`~pyaml.arrays.element_array.ElementArray` of matches + (possibly empty). + + Returns + ------- + T or ElementArray + The stored element for an exact literal name, otherwise an + ElementArray of matches. + + Raises + ------ + PyAMLException + If an exact literal name, or a literal entry within a list or + tuple, does not match any element in this holder, or a ``re:`` + pattern is not a valid regular expression. + """ + if isinstance(key, str) and not key.startswith(("re:", "~")) and not is_wildcard(key): + return self._peer._get(self._what, key, self._store) + names = resolve_names(self._store.keys(), key, what=self._what) + return ElementArray("", [self._store[n] for n in names]) + def add(self, m: T): """ Add an element to the holder's name-indexed store. diff --git a/pyaml/common/holders/rf_holder.py b/pyaml/common/holders/rf_holder.py index 257c396a..8b81d7c9 100644 --- a/pyaml/common/holders/rf_holder.py +++ b/pyaml/common/holders/rf_holder.py @@ -2,10 +2,12 @@ from typing import TYPE_CHECKING +from ...arrays.element_array import ElementArray from ...rf.rf_plant import RFPlant from ...rf.rf_transmitter import RFTransmitter from ..abstract import ReadWriteFloatScalar from ..element import __pyaml_repr__ +from ..name_matching import is_wildcard, resolve_names if TYPE_CHECKING: from .element_holder import ElementHolder @@ -50,6 +52,37 @@ def get(self, name: str) -> RFTransmitter: """ return self._peer._get("RFTransmitter", name, self._peer._RFTRANSMITTER) + def __getitem__(self, key: str | list[str] | tuple[str, ...]) -> "RFTransmitter | ElementArray": + """ + Return a transmitter, or a selection typed as a generic ElementArray. + + Parameters + ---------- + key : str, list[str] or tuple[str, ...] + An exact literal name returns the stored transmitter. An fnmatch + wildcard, a ``re:``-prefixed regular expression, or a list/tuple + of such patterns returns an ElementArray of matches (possibly + empty). + + Returns + ------- + RFTransmitter or ElementArray + The stored transmitter for an exact literal name, otherwise an + ElementArray of matches. + + Raises + ------ + PyAMLException + If an exact literal name, or a literal entry within a list or + tuple, does not match any transmitter, or a ``re:`` pattern is + not a valid regular expression. + """ + store = self._peer._RFTRANSMITTER + if isinstance(key, str) and not key.startswith(("re:", "~")) and not is_wildcard(key): + return self._peer._get("RFTransmitter", key, store) + names = resolve_names(store.keys(), key, what="RFTransmitter") + return ElementArray("", [store[n] for n in names]) + def add(self, rf: RFTransmitter): """ Add an RF transmitter to the holder. @@ -134,6 +167,37 @@ def get(self, name: str) -> RFPlant: """ return self._peer._get("RFPlant", name, self._peer._RFPLANT) + def __getitem__(self, key: str | list[str] | tuple[str, ...]) -> "RFPlant | ElementArray": + """ + Return an RF plant, or a selection typed as a generic ElementArray. + + Parameters + ---------- + key : str, list[str] or tuple[str, ...] + An exact literal name returns the stored RF plant. An fnmatch + wildcard, a ``re:``-prefixed regular expression, or a list/tuple + of such patterns returns an ElementArray of matches (possibly + empty). + + Returns + ------- + RFPlant or ElementArray + The stored RF plant for an exact literal name, otherwise an + ElementArray of matches. + + Raises + ------ + PyAMLException + If an exact literal name, or a literal entry within a list or + tuple, does not match any RF plant, or a ``re:`` pattern is not + a valid regular expression. + """ + store = self._peer._RFPLANT + if isinstance(key, str) and not key.startswith(("re:", "~")) and not is_wildcard(key): + return self._peer._get("RFPlant", key, store) + names = resolve_names(store.keys(), key, what="RFPlant") + return ElementArray("", [store[n] for n in names]) + def add(self, rf: RFPlant): """ Add an RF plant to the holder. diff --git a/pyaml/common/holders/tool_holder.py b/pyaml/common/holders/tool_holder.py index be7d17a0..d02c044b 100644 --- a/pyaml/common/holders/tool_holder.py +++ b/pyaml/common/holders/tool_holder.py @@ -5,6 +5,7 @@ from ...arrays.element_array import ElementArray from ..element import Element, __pyaml_repr__ from ..exception import PyAMLException +from ..name_matching import is_wildcard, resolve_names if TYPE_CHECKING: from ...tuning_tools.chromaticity import Chromaticity @@ -86,6 +87,37 @@ def get(self, name: str = None) -> "Element | ElementArray": return ElementArray("", list(self._peer._TOOLS.values())) return self._peer._get_tool(name) + def __getitem__(self, key: str | list[str] | tuple[str, ...]) -> "Element | ElementArray": + """ + Return a tool, or a selection typed as a generic ElementArray. + + Parameters + ---------- + key : str, list[str] or tuple[str, ...] + An exact literal name returns the stored tool. An fnmatch + wildcard, a ``re:``-prefixed regular expression, or a list/tuple + of such patterns returns an ElementArray of matches (possibly + empty). + + Returns + ------- + Element or ElementArray + The stored tool for an exact literal name, otherwise an + ElementArray of matches. + + Raises + ------ + PyAMLException + If an exact literal name, or a literal entry within a list or + tuple, does not match any tool, or a ``re:`` pattern is not a + valid regular expression. + """ + store = self._peer._TOOLS + if isinstance(key, str) and not key.startswith(("re:", "~")) and not is_wildcard(key): + return self._peer._get_tool(key) + names = resolve_names(store.keys(), key, what="Tool") + return ElementArray("", [store[n] for n in names]) + def _validate_type(self, name: str, obj: Element, expected_type: type) -> Element: """ Ensure a resolved default tool matches the type its accessor expects. diff --git a/pyaml/common/name_matching.py b/pyaml/common/name_matching.py new file mode 100644 index 00000000..9125cd3a --- /dev/null +++ b/pyaml/common/name_matching.py @@ -0,0 +1,84 @@ +"""Shared name and wildcard resolution convention used by every holder and array.""" + +import fnmatch +import re +from typing import Iterable + +from .exception import PyAMLException + +WILDCARD_CHARS = frozenset("*?[") + + +def is_wildcard(pattern: str) -> bool: + """Return True if `pattern` contains an fnmatch wildcard character.""" + return any(c in pattern for c in WILDCARD_CHARS) + + +def resolve_names(pool: Iterable[str], pattern: str | list[str] | tuple[str, ...], what: str = "Element") -> list[str]: + """ + Resolve one pattern, or several, against a pool of names. + + Parameters + ---------- + pool : Iterable[str] + Names to match against. + pattern : str, list[str] or tuple[str, ...] + A literal name, an fnmatch wildcard (triggered by ``*``, ``?`` or + ``[`` anywhere in the string), or a regular expression prefixed with + ``re:``. Prefix any of those with ``~`` to exclude its matches + instead of including them, mirroring the ``elements:`` selector list + convention used in YAML configuration. A sequence of patterns is + resolved entry by entry: non-``~`` entries are unioned, + de-duplicated, in first-encounter order, then anything matched by a + ``~`` entry is removed from that union. If every entry is + ``~``-prefixed (including a single lone ``~pattern``, treated as a + one-entry sequence), there is no explicit inclusion, so the base set + defaults to the full pool: ``~pattern`` alone means "everything + except pattern". + + Returns + ------- + list[str] + Matching names. + + Raises + ------ + PyAMLException + If a literal pattern (or the remainder of a ``~``-prefixed one) does + not match any name in `pool`, or if a `re:`-prefixed pattern is not + a valid regular expression. + """ + names = list(pool) + patterns = pattern if isinstance(pattern, (list, tuple)) else [pattern] + + included: list[str] = [] + seen: set[str] = set() + excluded: set[str] = set() + has_inclusion = False + for p in patterns: + if p.startswith("~"): + excluded.update(_resolve_one(names, p[1:], what)) + continue + has_inclusion = True + for name in _resolve_one(names, p, what): + if name not in seen: + seen.add(name) + included.append(name) + + base = included if has_inclusion else names + return [n for n in base if n not in excluded] + + +def _resolve_one(names: list[str], pattern: str, what: str) -> list[str]: + if pattern.startswith("re:"): + source = pattern[3:] + try: + compiled = re.compile(source) + except re.error as exc: + raise PyAMLException(f"Invalid regex '{source}': {exc}") from exc + return [n for n in names if compiled.fullmatch(n)] + if is_wildcard(pattern): + return [n for n in names if fnmatch.fnmatchcase(n, pattern)] + if pattern not in names: + raise PyAMLException(f"{what} {pattern} not defined") + return [pattern] diff --git a/tests/arrays/test_array_selection_types.py b/tests/arrays/test_array_selection_types.py index d5cc782b..11e65181 100644 --- a/tests/arrays/test_array_selection_types.py +++ b/tests/arrays/test_array_selection_types.py @@ -2,6 +2,7 @@ from pyaml.arrays.element_array import ElementArray from pyaml.arrays.magnet_array import MagnetArray +from pyaml.common.exception import PyAMLException @pytest.fixture @@ -60,3 +61,41 @@ def test_empty_selections_and_integer_indexing_keep_their_behavior(design): assert magnets["UNKNOWN*"] == [] assert magnets[0] is design.magnet.get(magnets.names()[0]) assert type(magnets - magnets) is list + + +def test_literal_name_miss_raises(design): + magnets = design.magnets.get() + + with pytest.raises(PyAMLException): + magnets["UNKNOWN"] + with pytest.raises(PyAMLException): + design.magnets["UNKNOWN"] + + +def test_literal_name_hit_still_returns_a_single_element_array(design): + magnets = design.magnets.get() + + selected = magnets["SH1A-C01-H"] + assert type(selected) is MagnetArray + assert selected.names() == ["SH1A-C01-H"] + + +def test_bracket_only_pattern_is_a_wildcard(design): + magnets = design.magnets.get() + assert sorted(magnets["SH1A-C0[12]-H"].names()) == ["SH1A-C01-H", "SH1A-C02-H"] + + +def test_regex_pattern_is_supported(design): + magnets = design.magnets.get() + assert sorted(magnets["re:^SH1A-C0[12]-H$"].names()) == ["SH1A-C01-H", "SH1A-C02-H"] + + +def test_list_of_patterns_is_supported(design): + magnets = design.magnets.get() + selected = magnets[["SH1A-C01-H", "SH1A-C02-H"]] + assert selected.names() == ["SH1A-C01-H", "SH1A-C02-H"] + + +def test_wildcard_matching_is_case_sensitive(design): + magnets = design.magnets.get() + assert magnets["sh1a*"] == [] diff --git a/tests/common/test_array_holder_navigation.py b/tests/common/test_array_holder_navigation.py index 7e630268..04cc2480 100644 --- a/tests/common/test_array_holder_navigation.py +++ b/tests/common/test_array_holder_navigation.py @@ -2,6 +2,7 @@ from pyaml.accelerator import Accelerator from pyaml.arrays.cfm_magnet_array import CombinedFunctionMagnetArray +from pyaml.common.exception import PyAMLException @pytest.fixture @@ -97,3 +98,76 @@ def test_issue_373_example(): one_magnet = sr.design.magnet.get("QF1E-C04") one_magnet.strength.set(0.8) assert one_magnet.strength.get() == pytest.approx(0.8) + + +def test_configured_exclusion_family_reproduced_with_selection_and_difference(): + """QForTest is QForTune minus a `~name`/`~pattern` YAML exclusion; `[]` and `-` reproduce it.""" + sr = Accelerator.load("tests/config/EBSTune-patterns.yaml", ignore_external=True) + sr.design.get_lattice().disable_6d() + + q_for_tune = sr.design.magnets.get("QForTune") + q_for_test = sr.design.magnets.get("QForTest") # QForTune, minus ~QF1E-C05 and ~Q???-C06 + + reproduced = q_for_tune - sr.design.magnets["QF1E-C05"] - sr.design.magnets["Q???-C06"] + + assert reproduced == q_for_test + + +def test_configured_exclusion_family_reproduced_in_a_single_getitem_call(): + """The YAML `elements:` selector list (`[QD2*, QF1*, ~QF1E-C05, ~Q???-C06]`) works verbatim through `[]`. + + Selection goes through the top-level holder, not `.magnets`, so it resolves patterns + against the same global element pool `find_elements()`/`_fill_array` use to build + `QForTest` in the YAML, and so lands in the same order. + """ + sr = Accelerator.load("tests/config/EBSTune-patterns.yaml", ignore_external=True) + sr.design.get_lattice().disable_6d() + + q_for_test = sr.design.magnets.get("QForTest") + reproduced = sr.design[["QD2*", "QF1*", "~QF1E-C05", "~Q???-C06"]] + + assert type(reproduced) is type(q_for_test) + assert reproduced == q_for_test + + +def test_magnet_holder_getitem_exact_name_matches_get(holder): + assert holder.magnet["SH1A-C01-H"] is holder.magnet.get("SH1A-C01-H") + + +def test_magnet_holder_getitem_exact_name_miss_raises(holder): + with pytest.raises(PyAMLException): + holder.magnet["UNKNOWN"] + + +def test_magnet_holder_getitem_wildcard_and_list_return_an_array(holder): + wildcard = holder.magnet["SH1A-C01-[HV]"] + assert sorted(wildcard.names()) == ["SH1A-C01-H", "SH1A-C01-V"] + assert holder.magnet["MISSING*"].names() == [] + + listed = holder.magnet[["SH1A-C01-H", "SH1A-C01-V"]] + assert listed.names() == ["SH1A-C01-H", "SH1A-C01-V"] + + +def test_magnet_holder_getitem_regex(holder): + matching = holder.magnet["re:^SH1A-C0[12]-H$"] + assert sorted(matching.names()) == ["SH1A-C01-H", "SH1A-C02-H"] + + +def test_combined_function_magnet_holder_getitem_exact_name_matches_get(holder): + assert holder.combined_function_magnet["SH1A-C01"] is holder.combined_function_magnet.get("SH1A-C01") + + +def test_combined_function_magnet_holder_getitem_exact_name_miss_raises(holder): + with pytest.raises(PyAMLException): + holder.combined_function_magnet["UNKNOWN"] + + +def test_serialized_magnet_holder_getitem_exact_name_matches_get(serialized_holder): + assert serialized_holder.serialized_magnet["mySeriesOfMagnets"] is serialized_holder.serialized_magnet.get( + "mySeriesOfMagnets" + ) + + +def test_serialized_magnet_holder_getitem_exact_name_miss_raises(serialized_holder): + with pytest.raises(PyAMLException): + serialized_holder.serialized_magnet["UNKNOWN"] diff --git a/tests/common/test_element_holder_collection.py b/tests/common/test_element_holder_collection.py index 4a55fadf..c6c889d5 100644 --- a/tests/common/test_element_holder_collection.py +++ b/tests/common/test_element_holder_collection.py @@ -15,9 +15,10 @@ def holder(accelerator_from_fragments, sr_configuration_fragments): return sr.design -def test_exact_name_returns_the_element_or_none(holder): +def test_exact_name_returns_the_element_or_raises(holder): assert holder["BPM_C04-01"] is holder.diagnostic.bpm.get("BPM_C04-01") - assert holder["UNKNOWN"] is None + with pytest.raises(PyAMLException): + holder["UNKNOWN"] def test_get_and_full_slice_keep_registration_order(holder): @@ -113,7 +114,8 @@ def test_empty_holder_returns_empty_collections(ebs_lattice_file): assert type(holder.get()) is ElementArray assert holder[:].names() == [] assert holder["BPM*"].names() == [] - assert holder["BPM_C04-01"] is None + with pytest.raises(PyAMLException): + holder["BPM_C04-01"] def test_invalid_indices_and_keys_raise_clear_errors(holder): @@ -148,3 +150,128 @@ def test_existing_array_field_filters_still_work(holder): def test_existing_empty_intersection_stays_a_list(holder): assert type(holder["SH1A-C0?-H"] & holder["SH1A-C0?-V"]) is list + + +def test_intersecting_two_wildcard_selections(holder): + """Two overlapping wildcard selections combined with `&` narrow down to their overlap.""" + by_cell = holder["SH1A-C01*"] # SH1A-C01-H, SH1A-C01-V, SH1A-C01-SQ + by_plane = holder["SH1A-C0?-H"] # SH1A-C01-H, SH1A-C02-H + + selected = by_cell & by_plane + + assert isinstance(selected, MagnetArray) + assert selected.names() == ["SH1A-C01-H"] + + +def test_union_of_a_list_selection_and_a_wildcard_selection(holder): + """A list-of-patterns selection and a wildcard selection combine with `|` like any array.""" + cell_01 = holder[["SH1A-C01-H", "SH1A-C01-V"]] + cell_02 = holder["SH1A-C02*"] # the CFM magnet SH1A-C02, plus its -H, -V, -SQ sub-magnets + + combined = cell_01 | cell_02 + + # Mixes a CombinedFunctionMagnet (SH1A-C02) with plain Magnets, so the union + # stays a generic ElementArray rather than a MagnetArray. + assert type(combined) is ElementArray + assert sorted(combined.names()) == [ + "SH1A-C01-H", + "SH1A-C01-V", + "SH1A-C02", + "SH1A-C02-H", + "SH1A-C02-SQ", + "SH1A-C02-V", + ] + + +def test_difference_with_a_regex_selection(holder): + """`-` removes a `re:` selection's matches from a wildcard selection, like any array.""" + both_planes_h = holder["SH1A-C0?-H"] # SH1A-C01-H, SH1A-C02-H + cell_01_h = holder["re:^SH1A-C01-H$"] + + selected = both_planes_h - cell_01_h + + assert isinstance(selected, MagnetArray) + assert selected.names() == ["SH1A-C02-H"] + + +def test_disjoint_regex_and_list_selections_produce_a_plain_list(holder): + """The same empty-result-is-a-plain-list quirk applies to the new selection styles too.""" + cell_01_h = holder["re:^SH1A-C01-H$"] + cell_02_h_v = holder[["SH1A-C02-H", "SH1A-C02-V"]] + + assert type(cell_01_h & cell_02_h_v) is list + + +def test_selection_combined_with_a_configured_family_via_regex_and_list(holder): + """`re:` and list-of-patterns selections intersect with a configured family, same as wildcards.""" + el_array = holder.get_elements("ElArray") # BPM_C04-01, BPM_C04-02, SH1A-C01-V, SH1A-C02-H + + by_regex = holder["re:^SH1A-C0[12]-H$"] & el_array + assert isinstance(by_regex, MagnetArray) + assert by_regex.names() == ["SH1A-C02-H"] + + by_list = holder[["BPM_C04-01", "BPM_C04-02"]] & el_array + assert isinstance(by_list, BPMArray) + assert by_list.names() == ["BPM_C04-01", "BPM_C04-02"] + + +def test_regex_pattern_is_supported(holder): + assert holder["re:^BPM_C04-0[12]$"].names() == ["BPM_C04-01", "BPM_C04-02"] + assert holder["re:^NOTHING$"].names() == [] + + +def test_invalid_regex_raises_pyaml_exception(holder): + with pytest.raises(PyAMLException, match="Invalid regex"): + holder["re:("] + + +def test_list_of_patterns_unions_results(holder): + selected = holder[["BPM_C04-01", "SH1A-C0?-H"]] + assert selected.names() == ["BPM_C04-01", "SH1A-C01-H", "SH1A-C02-H"] + + +def test_list_of_patterns_with_missing_literal_raises(holder): + with pytest.raises(PyAMLException): + holder[["BPM_C04-01", "UNKNOWN"]] + + +def test_wildcard_matching_is_case_sensitive(holder): + assert holder["bpm_c04*"].names() == [] + + +def test_find_elements_literal_miss_raises(holder): + with pytest.raises(PyAMLException): + holder.find_elements("UNKNOWN") + + +def test_find_elements_bracket_only_pattern_is_a_wildcard(holder): + assert holder.find_elements("BPM_C04-0[12]") == ["BPM_C04-01", "BPM_C04-02"] + + +def test_find_elements_supports_regex(holder): + assert holder.find_elements("re:^BPM_C04-0[12]$") == ["BPM_C04-01", "BPM_C04-02"] + + +def test_find_elements_supports_a_list_of_patterns(holder): + assert holder.find_elements(["BPM_C04-01", "SH1A-C0?-H"]) == ["BPM_C04-01", "SH1A-C01-H", "SH1A-C02-H"] + + +def test_find_elements_supports_an_exclusion_in_a_list(holder): + both_planes = holder.find_elements(["SH1A-C0?-H", "SH1A-C0?-V"]) + assert both_planes == ["SH1A-C01-H", "SH1A-C02-H", "SH1A-C01-V", "SH1A-C02-V"] + + minus_one = holder.find_elements(["SH1A-C0?-H", "SH1A-C0?-V", "~SH1A-C02-V"]) + assert minus_one == ["SH1A-C01-H", "SH1A-C02-H", "SH1A-C01-V"] + + +def test_getitem_supports_an_exclusion_in_a_list(holder): + selected = holder[["SH1A-C0?-H", "SH1A-C0?-V", "~SH1A-C02-V"]] + assert selected.names() == ["SH1A-C01-H", "SH1A-C02-H", "SH1A-C01-V"] + + +def test_getitem_lone_exclusion_pattern_means_everything_except(holder): + all_names = holder.get_all_elements() + selected = holder["~QF1A-C01"] + + assert "QF1A-C01" not in selected.names() + assert len(selected) == len(all_names) - 1 diff --git a/tests/common/test_name_matching.py b/tests/common/test_name_matching.py new file mode 100644 index 00000000..c7365862 --- /dev/null +++ b/tests/common/test_name_matching.py @@ -0,0 +1,114 @@ +import pytest + +from pyaml.common.exception import PyAMLException +from pyaml.common.name_matching import is_wildcard, resolve_names + +NAMES = ["BPM_C04-01", "BPM_C04-02", "QF1A-C01", "QF1A-C02"] + + +def test_is_wildcard_triggers_on_star_question_and_bracket(): + assert is_wildcard("BPM*") + assert is_wildcard("BPM?") + assert is_wildcard("BPM[01]") + assert not is_wildcard("BPM_C04-01") + assert not is_wildcard("re:^BPM") + + +def test_literal_hit_returns_the_single_name(): + assert resolve_names(NAMES, "QF1A-C01") == ["QF1A-C01"] + + +def test_literal_miss_raises(): + with pytest.raises(PyAMLException, match="Element UNKNOWN not defined"): + resolve_names(NAMES, "UNKNOWN") + + +def test_literal_miss_uses_the_provided_what_label(): + with pytest.raises(PyAMLException, match="Magnet UNKNOWN not defined"): + resolve_names(NAMES, "UNKNOWN", what="Magnet") + + +def test_wildcard_hit_and_empty_are_not_errors(): + assert resolve_names(NAMES, "BPM*") == ["BPM_C04-01", "BPM_C04-02"] + assert resolve_names(NAMES, "MISSING*") == [] + + +def test_wildcard_is_case_sensitive(): + assert resolve_names(NAMES, "bpm*") == [] + + +def test_bracket_only_pattern_is_a_wildcard_not_a_literal(): + assert resolve_names(NAMES, "QF1A-C0[12]") == ["QF1A-C01", "QF1A-C02"] + + +def test_regex_hit_and_empty_are_not_errors(): + assert resolve_names(NAMES, "re:^BPM_C04-0[12]$") == ["BPM_C04-01", "BPM_C04-02"] + assert resolve_names(NAMES, "re:^NOTHING$") == [] + + +def test_regex_is_a_fullmatch(): + assert resolve_names(NAMES, "re:BPM_C04") == [] + + +def test_invalid_regex_raises_pyaml_exception_not_re_error(): + with pytest.raises(PyAMLException, match="Invalid regex"): + resolve_names(NAMES, "re:(") + + +def test_list_of_patterns_unions_and_deduplicates_in_order(): + result = resolve_names(NAMES, ["QF1A-C01", "BPM*", "QF1A-C01"]) + assert result == ["QF1A-C01", "BPM_C04-01", "BPM_C04-02"] + + +def test_tuple_of_patterns_is_also_accepted(): + assert resolve_names(NAMES, ("QF1A-C01", "QF1A-C02")) == ["QF1A-C01", "QF1A-C02"] + + +def test_list_with_one_missing_literal_raises(): + with pytest.raises(PyAMLException, match="Element UNKNOWN not defined"): + resolve_names(NAMES, ["QF1A-C01", "UNKNOWN"]) + + +def test_list_mixing_literal_wildcard_and_regex(): + result = resolve_names(NAMES, ["QF1A-C01", "BPM*", "re:^QF1A-C02$"]) + assert result == ["QF1A-C01", "BPM_C04-01", "BPM_C04-02", "QF1A-C02"] + + +def test_overlapping_wildcard_patterns_in_a_list_do_not_duplicate_a_match(): + names = ["BPM_01", "BPM_02", "BPM_03", "BPM_04", "BPM_05"] + result = resolve_names(names, ["BPM_0[1-3]", "BPM_0[3-5]"]) + assert result == ["BPM_01", "BPM_02", "BPM_03", "BPM_04", "BPM_05"] + + +def test_lone_exclusion_pattern_means_everything_except(): + assert resolve_names(NAMES, "~QF1A-C01") == ["BPM_C04-01", "BPM_C04-02", "QF1A-C02"] + + +def test_list_of_only_exclusion_patterns_means_everything_except(): + result = resolve_names(NAMES, ["~QF1A-C01", "~QF1A-C02"]) + assert result == ["BPM_C04-01", "BPM_C04-02"] + + +def test_exclusion_removes_matches_from_an_inclusion_pattern(): + assert resolve_names(NAMES, ["BPM*", "~BPM_C04-01"]) == ["BPM_C04-02"] + + +def test_exclusion_order_does_not_matter(): + forward = resolve_names(NAMES, ["BPM*", "~BPM_C04-01"]) + backward = resolve_names(NAMES, ["~BPM_C04-01", "BPM*"]) + assert forward == backward == ["BPM_C04-02"] + + +def test_exclusion_of_a_wildcard_pattern_removes_all_its_matches(): + result = resolve_names(NAMES, ["QF1A-C01", "QF1A-C02", "BPM_C04-01", "~QF1A-*"]) + assert result == ["BPM_C04-01"] + + +def test_exclusion_of_a_regex_pattern_removes_all_its_matches(): + result = resolve_names(NAMES, ["BPM*", "~re:^BPM_C04-01$"]) + assert result == ["BPM_C04-02"] + + +def test_exclusion_of_a_missing_literal_raises(): + with pytest.raises(PyAMLException, match="Element UNKNOWN not defined"): + resolve_names(NAMES, ["BPM*", "~UNKNOWN"]) diff --git a/tests/diagnostics/test_diagnostic_accessors.py b/tests/diagnostics/test_diagnostic_accessors.py index 2247db71..342c04dd 100644 --- a/tests/diagnostics/test_diagnostic_accessors.py +++ b/tests/diagnostics/test_diagnostic_accessors.py @@ -79,3 +79,71 @@ def test_diagnostic_bpms_returns_named_array(): bpms = design.diagnostic.bpms.get("BPM") assert design.diagnostic.bpm.get("BPM_C04-04") in bpms + + +def test_diagnostic_getitem_exact_name_matches_get(): + design = Accelerator.load( + "tests/config/EBSOrbit.yaml", + ignore_external=True, + include_locations=False, + ).design + + assert design.diagnostic["BETATRON_TUNE"] is design.diagnostic.get("BETATRON_TUNE") + + +def test_diagnostic_getitem_exact_name_miss_raises(): + design = Accelerator.load( + "tests/config/EBSOrbit.yaml", + ignore_external=True, + include_locations=False, + ).design + + with pytest.raises(PyAMLException): + design.diagnostic["UNKNOWN"] + + +def test_diagnostic_getitem_wildcard_returns_an_array(): + design = Accelerator.load( + "tests/config/EBSOrbit.yaml", + ignore_external=True, + include_locations=False, + ).design + + matching = design.diagnostic["BETATRON*"] + assert matching.names() == ["BETATRON_TUNE"] + assert design.diagnostic["MISSING*"].names() == [] + + +def test_diagnostic_bpm_getitem_exact_name_matches_get(): + design = Accelerator.load( + "tests/config/EBSOrbit.yaml", + ignore_external=True, + include_locations=False, + ).design + + assert design.diagnostic.bpm["BPM_C04-04"] is design.diagnostic.bpm.get("BPM_C04-04") + + +def test_diagnostic_bpm_getitem_exact_name_miss_raises(): + design = Accelerator.load( + "tests/config/EBSOrbit.yaml", + ignore_external=True, + include_locations=False, + ).design + + with pytest.raises(PyAMLException): + design.diagnostic.bpm["UNKNOWN"] + + +def test_diagnostic_bpm_getitem_wildcard_and_list_return_an_array(): + design = Accelerator.load( + "tests/config/EBSOrbit.yaml", + ignore_external=True, + include_locations=False, + ).design + + wildcard = design.diagnostic.bpm["BPM_C04-0[14]"] + assert sorted(wildcard.names()) == ["BPM_C04-01", "BPM_C04-04"] + + listed = design.diagnostic.bpm[["BPM_C04-01", "BPM_C04-04"]] + assert listed.names() == ["BPM_C04-01", "BPM_C04-04"] diff --git a/tests/rf/test_rf_holder.py b/tests/rf/test_rf_holder.py new file mode 100644 index 00000000..784577a8 --- /dev/null +++ b/tests/rf/test_rf_holder.py @@ -0,0 +1,48 @@ +import pytest + +from pyaml.accelerator import Accelerator +from pyaml.common.exception import PyAMLException + + +@pytest.fixture +def design(): + return Accelerator.load("tests/config/EBS_rf_multi.yaml", ignore_external=True).design + + +def test_rf_holder_getitem_exact_name_matches_get(design): + assert design.rf["DEFAULT_RF_PLANT"] is design.rf.get("DEFAULT_RF_PLANT") + + +def test_rf_holder_getitem_exact_name_miss_raises(design): + with pytest.raises(PyAMLException): + design.rf["UNKNOWN"] + + +def test_rf_holder_getitem_wildcard_returns_an_array(design): + matching = design.rf["DEFAULT_*"] + assert matching.names() == ["DEFAULT_RF_PLANT"] + assert design.rf["MISSING*"].names() == [] + + +def test_rf_transmitter_holder_getitem_exact_name_matches_get(design): + assert design.rf.transmitter["RFTRA1"] is design.rf.transmitter.get("RFTRA1") + + +def test_rf_transmitter_holder_getitem_exact_name_miss_raises(design): + with pytest.raises(PyAMLException): + design.rf.transmitter["UNKNOWN"] + + +def test_rf_transmitter_holder_getitem_wildcard_returns_an_array(design): + matching = design.rf.transmitter["RFTRA*"] + assert sorted(matching.names()) == ["RFTRA1", "RFTRA2", "RFTRA_HARMONIC"] + + +def test_rf_transmitter_holder_getitem_list_of_patterns(design): + selected = design.rf.transmitter[["RFTRA1", "RFTRA2"]] + assert selected.names() == ["RFTRA1", "RFTRA2"] + + +def test_rf_transmitter_holder_getitem_regex(design): + matching = design.rf.transmitter["re:^RFTRA[12]$"] + assert sorted(matching.names()) == ["RFTRA1", "RFTRA2"] diff --git a/tests/tuning_tools/test_tool_holder.py b/tests/tuning_tools/test_tool_holder.py new file mode 100644 index 00000000..3821fcf3 --- /dev/null +++ b/tests/tuning_tools/test_tool_holder.py @@ -0,0 +1,38 @@ +import pytest + +from pyaml.accelerator import Accelerator +from pyaml.common.exception import PyAMLException + + +@pytest.fixture +def design(): + return Accelerator.load( + "tests/config/EBSOrbit.yaml", + ignore_external=True, + include_locations=False, + ).design + + +def test_tool_holder_getitem_exact_name_matches_get(design): + assert design.tool["DEFAULT_TUNE_CORRECTION"] is design.tool.get("DEFAULT_TUNE_CORRECTION") + + +def test_tool_holder_getitem_exact_name_miss_raises(design): + with pytest.raises(PyAMLException): + design.tool["UNKNOWN"] + + +def test_tool_holder_getitem_wildcard_returns_an_array(design): + matching = design.tool["DEFAULT_TUNE*"] + assert sorted(matching.names()) == ["DEFAULT_TUNE_CORRECTION", "DEFAULT_TUNE_RESPONSE_MATRIX"] + assert design.tool["MISSING*"].names() == [] + + +def test_tool_holder_getitem_list_of_patterns(design): + selected = design.tool[["DEFAULT_TUNE_CORRECTION", "DEFAULT_ORBIT_CORRECTION"]] + assert selected.names() == ["DEFAULT_TUNE_CORRECTION", "DEFAULT_ORBIT_CORRECTION"] + + +def test_tool_holder_getitem_regex(design): + matching = design.tool["re:^DEFAULT_(TUNE|ORBIT)_CORRECTION$"] + assert sorted(matching.names()) == ["DEFAULT_ORBIT_CORRECTION", "DEFAULT_TUNE_CORRECTION"]