diff options
Diffstat (limited to 'python/src/types/selection.py')
| -rw-r--r-- | python/src/types/selection.py | 153 |
1 files changed, 0 insertions, 153 deletions
diff --git a/python/src/types/selection.py b/python/src/types/selection.py deleted file mode 100644 index d75c9f0..0000000 --- a/python/src/types/selection.py +++ /dev/null @@ -1,153 +0,0 @@ -from dataclasses import dataclass, field -from typing import Callable - -from .buffer import Buffer - - -@dataclass -class Position: - pos: int - - -@dataclass -class Range: - start: Position - end: Position - - -@dataclass -class Ranges: - ranges: list[Range] - - -@dataclass -class MultiRanges: - multi_ranges: dict[Buffer, Ranges] - - -@dataclass -class Selection: - selection: Position | Range | Ranges | MultiRanges - buffers: set[Buffer] = field(default_factory=set) - - def __post_init__(self): - if type(self.selection) is MultiRanges: - self.buffers = set(self.selection.multi_ranges) - else: - assert len(self.buffers) == 1 - - @classmethod - def empty(cls) -> "Selection": - return cls(selection=Ranges(ranges=[]), buffers=set()) - - def promote(self, to_rank: int) -> "Selection": - if to_rank < 0 or to_rank > 3: - raise ValueError(f"Invalid rank: {to_rank}") - if self.rank > to_rank: - raise ValueError(f"Cannot coerce selection of rank {self.rank} to rank {to_rank}") - if self.rank == to_rank: - return self - match self.selection: - case Position(_): - as_range = Range( - start=Position(self.selection.pos), - end=Position(self.selection.pos + 1), - ) - as_ranges = Ranges(ranges=[as_range]) - if to_rank == 1: - return Selection(selection=as_range, buffers=self.buffers) - if to_rank == 2: - return Selection(selection=as_ranges, buffers=self.buffers) - if to_rank == 3: - return Selection(MultiRanges({buffer: as_ranges for buffer in self.buffers})) - case Range(_, _): - as_ranges = Ranges(ranges=[self.selection]) - if to_rank == 2: - return Selection(selection=as_ranges, buffers=self.buffers) - if to_rank == 3: - return Selection(MultiRanges({buffer: as_ranges for buffer in self.buffers})) - case Ranges(_): - return Selection(MultiRanges({buffer: self.selection for buffer in self.buffers})) - - raise RuntimeError("Unexpected promotion issue") - - @classmethod - def union(cls, *selections: "Selection") -> "Selection": - if not selections: - return cls.empty() - - all_buffers = set.union(*[s.buffers for s in selections]) - - if len(all_buffers) == 0: - return cls.empty() - - if len(all_buffers) > 1: - max_rank = 3 - else: - max_rank = max(s.rank for s in selections) - if max_rank < 2: - max_rank = 2 - - promoted = [s.promote(max_rank) for s in selections] - - if max_rank == 2: - all_ranges = [] - for sel in promoted: - if type(sel.selection) is not Ranges: - raise ValueError( - f"Selection of rank {sel.rank} is not a range. Should have been promoted to rank 2." - ) - all_ranges.extend(sel.selection.ranges) - return cls(selection=Ranges(all_ranges), buffers=all_buffers) - - if max_rank == 3: - buffer_ranges: dict[Buffer, list[Range]] = {buf: [] for buf in all_buffers} - for sel in promoted: - if type(sel.selection) is not MultiRanges: - raise ValueError( - f"Selection of rank {sel.rank} is not a multi-range. Should have been promoted to rank 3." - ) - for buffer, ranges in sel.selection.multi_ranges.items(): - buffer_ranges[buffer].extend(ranges.ranges) - - multi = MultiRanges( - {buffer: Ranges(ranges) for buffer, ranges in buffer_ranges.items()} - ) - return cls(selection=multi) - - raise RuntimeError("Unexpected union issue") - - @property - def rank(self) -> int: - match self.selection: - case Position(_): - return 0 - case Range(_, _): - return 1 - case Ranges(_): - return 2 - case MultiRanges(_): - return 3 - - def broadcast( - self, - rank_zero: Callable[[Position, Buffer], "Selection"], - rank_one: Callable[[Range, Buffer], "Selection"], - ) -> "Selection": - match self.selection: - case Position(_): - return rank_zero(self.selection, list(self.buffers)[0]) - - case Range(_, _) as r: - return rank_one(r, list(self.buffers)[0]) - - case Ranges(ranges): - buffer = list(self.buffers)[0] - return Selection.union(*[rank_one(r, buffer) for r in ranges]) - - case MultiRanges(multi): - all_results = [] - for buffer, ranges in multi.items(): - for r in ranges.ranges: - all_results.append(rank_one(r, buffer)) - return Selection.union(*all_results) |
