diff --git a/pyaml/arrays/element_array.py b/pyaml/arrays/element_array.py index 56985b43..ea620c8e 100644 --- a/pyaml/arrays/element_array.py +++ b/pyaml/arrays/element_array.py @@ -182,6 +182,25 @@ def __auto_array(self, elements: list[Element]): if len(elements) == 0: return [] + return self._typed_array(elements) + + def _typed_array(self, elements: list[Element]) -> "ElementArray": + """Build a collection using the most specific compatible array type. + + Parameters + ---------- + elements : list[Element] + Selected references, in their desired order. + + Returns + ------- + ElementArray + Specialized array when possible, otherwise a generic array. + An empty selection returns an empty generic array. + """ + if not elements: + return self.__create_array("", Element, elements) + import inspect def mro_as_list(cls: type) -> list[type]: @@ -210,6 +229,22 @@ 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 --- diff --git a/pyaml/common/holders/element_holder.py b/pyaml/common/holders/element_holder.py index 1db1daa6..aa2b16dc 100644 --- a/pyaml/common/holders/element_holder.py +++ b/pyaml/common/holders/element_holder.py @@ -3,7 +3,7 @@ import fnmatch import re from abc import ABCMeta, abstractmethod -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, overload from ...arrays.element_array import ElementArray from ...bpm.bpm import BPM @@ -419,6 +419,95 @@ def _get(self, what, name, array) -> Element: return array[name] # Generic elements + def get(self) -> ElementArray: + """Return all registered elements in insertion order. + + Returns + ------- + ElementArray + New unnamed container sharing the registered element references. + + Notes + ----- + Registration order is not necessarily longitudinal lattice order. + Changing the returned container does not change the holder registry. + Each call reflects the current registry. + + Examples + -------- + >>> elements = sr.live.get() + >>> names = elements.names() + """ + return ElementArray("", list(self._ALL.values())) + + @overload + def __getitem__(self, key: int) -> Element: ... + + @overload + def __getitem__(self, key: slice) -> ElementArray: ... + + @overload + def __getitem__(self, key: str) -> Element | ElementArray | None: ... + + def __getitem__(self, key: int | slice | str) -> Element | ElementArray | None: + """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. + + 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. + The full slice ``[:]`` returns a generic ElementArray, like get(). + + Raises + ------ + IndexError + If the index is out of bounds. + TypeError + If the key is neither an integer, a slice, nor a string. + ValueError + If a slice has a zero step. + + Notes + ----- + 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. + + Examples + -------- + >>> bpm = sr.live["BPM01"] + >>> missing = sr.live["UNKNOWN"] # None + >>> 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 + >>> 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 isinstance(key, int): + return list(self._ALL.values())[key] + if isinstance(key, slice): + elements = self.get() + if key == slice(None): + return elements + return elements._typed_array(list(elements)[key]) + raise TypeError("ElementHolder keys must be integers, slices or strings") + def fill_element_array(self, arrayName: str, elementNames: list[str]): """ Create and register a generic element array. diff --git a/tests/common/test_element_holder_collection.py b/tests/common/test_element_holder_collection.py new file mode 100644 index 00000000..937d2da7 --- /dev/null +++ b/tests/common/test_element_holder_collection.py @@ -0,0 +1,150 @@ +import pytest + +from pyaml.arrays.bpm_array import BPMArray +from pyaml.arrays.element_array import ElementArray +from pyaml.arrays.magnet_array import MagnetArray +from pyaml.bpm.bpm import BPM +from pyaml.common.exception import PyAMLException +from pyaml.lattice.simulator import Simulator + + +@pytest.fixture +def holder(accelerator_from_fragments, sr_configuration_fragments): + sr = accelerator_from_fragments(*sr_configuration_fragments) + sr.design.get_lattice().disable_6d() + return sr.design + + +def test_exact_name_returns_the_element_or_none(holder): + assert holder["BPM_C04-01"] is holder.bpm.get("BPM_C04-01") + assert holder["UNKNOWN"] is None + + +def test_get_and_full_slice_keep_registration_order(holder): + names = [element.get_name() for element in holder.get_all_elements()] + + assert type(holder.get()) is ElementArray + assert type(holder[:]) is ElementArray + assert holder.get().names() == names + assert holder[:].names() == names + assert names != sorted(names) + + +def test_collections_are_independent_but_share_elements(holder): + collection = holder.get() + first = holder[0] + original_names = collection.names() + + assert collection[0] is first + collection.clear() + assert holder[0] is first + assert holder.get().names() == original_names + + +def test_new_calls_reflect_registry_additions(holder): + previous = holder.get() + holder.fill_device([BPM("EXTRA_BPM", lattice_names="list(BPM_C04-01)")]) + + assert "EXTRA_BPM" not in previous.names() + assert holder.get().names() == previous.names() + ["EXTRA_BPM"] + assert holder["EXTRA_BPM"] is holder.bpm.get("EXTRA_BPM") + + +def test_patterns_always_return_arrays(holder): + assert isinstance(holder["BPM*"], BPMArray) + assert holder["BPM*"].names() == ["BPM_C04-01", "BPM_C04-02"] + assert isinstance(holder["BPM_C04-0[1]"], BPMArray) + assert holder["BPM_C04-0[1]"].names() == ["BPM_C04-01"] + assert holder["BPM_C04-0?"].names() == ["BPM_C04-01", "BPM_C04-02"] + assert type(holder["MISSING*"]) is ElementArray + assert holder["MISSING*"].names() == [] + + +def test_character_classes_in_name_patterns(holder): + assert holder["BPM_C04-0[12]"].names() == ["BPM_C04-01", "BPM_C04-02"] + assert holder["BPM_C04-0[1-2]"].names() == ["BPM_C04-01", "BPM_C04-02"] + assert holder["SH1A-C01-[HV]*"].names() == ["SH1A-C01-H", "SH1A-C01-V"] + assert holder["BPM_C04-0[!2]"].names() == ["BPM_C04-01"] + + +def test_colons_are_part_of_names(holder): + holder.fill_device([BPM("CELL04:BPM01", lattice_names="list(BPM_C04-01)")]) + + assert holder["CELL04:BPM01"] is holder.bpm.get("CELL04:BPM01") + assert holder["CELL04:BPM*"].names() == ["CELL04:BPM01"] + assert holder["model_name:*"].names() == [] + + +def test_magnet_subclasses_share_a_typed_array_in_either_order(holder): + horizontal = holder.magnet.get("SH1A-C01-H") + vertical = holder.magnet.get("SH1A-C01-V") + start = holder.get_all_elements().index(horizontal) + + assert type(horizontal) is not type(vertical) + assert isinstance(holder["SH1A-C01-[HV]"], MagnetArray) + assert isinstance(holder[start : start + 2], MagnetArray) + assert isinstance(holder[start + 1 : start - 1 : -1], MagnetArray) + assert holder[start + 1 : start - 1 : -1].names() == ["SH1A-C01-V", "SH1A-C01-H"] + + +def test_mixed_selection_returns_a_generic_array(holder): + selected = holder["*-C01*"] + + assert type(selected) is ElementArray + assert "QF1A-C01" in selected.names() + assert "SH1A-C01" in selected.names() + + +def test_indices_and_slices_follow_insertion_order(holder): + registered = holder.get_all_elements() + + assert holder[0] is registered[0] + assert holder[-1] is registered[-1] + assert list(holder[1:3]) == registered[1:3] + assert list(holder[::2]) == registered[::2] + assert isinstance(holder[-2:], BPMArray) + assert holder[-2:].names() == ["BPM_C04-01", "BPM_C04-02"] + assert type(holder[len(registered) :]) is ElementArray + + +def test_empty_holder_returns_empty_collections(ebs_lattice_file): + holder = Simulator(name="empty", lattice=str(ebs_lattice_file)) + + assert type(holder.get()) is ElementArray + assert holder[:].names() == [] + assert holder["BPM*"].names() == [] + assert holder["BPM_C04-01"] is None + + +def test_invalid_indices_and_keys_raise_clear_errors(holder): + size = len(holder.get_all_elements()) + + with pytest.raises(IndexError): + holder[size] + with pytest.raises(IndexError): + holder[-size - 1] + with pytest.raises(TypeError): + holder[1.5] + with pytest.raises(ValueError): + holder[::0] + + +def test_selection_intersects_with_a_configured_family(holder): + selected = holder["SH1A-C0?-H"] & holder.get_elements("ElArray") + + assert isinstance(selected, MagnetArray) + assert selected.names() == ["SH1A-C02-H"] + assert holder.get_element("SH1A-C02-H") is holder["SH1A-C02-H"] + assert holder.get_all_elements() == list(holder.get()) + with pytest.raises(PyAMLException): + holder.get_element("UNKNOWN") + + +def test_existing_array_field_filters_still_work(holder): + selected = holder["SH1A-C0?-H"]["model_name:SH1A-C01"] + + assert selected.names() == ["SH1A-C01-H"] + + +def test_existing_empty_intersection_stays_a_list(holder): + assert type(holder["SH1A-C0?-H"] & holder["SH1A-C0?-V"]) is list diff --git a/tests/test_load_conf_with_code.py b/tests/test_load_conf_with_code.py index 42286ba2..de476102 100644 --- a/tests/test_load_conf_with_code.py +++ b/tests/test_load_conf_with_code.py @@ -11,3 +11,8 @@ def test_load_conf_with_code(): bpms = sr.live.bpms.get("BPM") assert bpms is not None assert len(bpms) == 320 + + assert sr.live[bpms[0].get_name()] is bpms[0] + assert sr.live["BPM*"].names() == bpms.names() + assert sr.live[:].names() == [element.get_name() for element in sr.live.get_all_elements()] + assert sr.design["BPM*"].names() == sr.design.bpms.get("BPM").names()