Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 35 additions & 0 deletions pyaml/arrays/element_array.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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 ---
Expand Down
91 changes: 90 additions & 1 deletion pyaml/common/holders/element_holder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
150 changes: 150 additions & 0 deletions tests/common/test_element_holder_collection.py
Original file line number Diff line number Diff line change
@@ -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
5 changes: 5 additions & 0 deletions tests/test_load_conf_with_code.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading