aboutsummaryrefslogtreecommitdiff
path: root/python/src/types/selection.py
diff options
context:
space:
mode:
Diffstat (limited to 'python/src/types/selection.py')
-rw-r--r--python/src/types/selection.py153
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)