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)