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.py161
1 files changed, 161 insertions, 0 deletions
diff --git a/python/src/types/selection.py b/python/src/types/selection.py
new file mode 100644
index 0000000..478b5b9
--- /dev/null
+++ b/python/src/types/selection.py
@@ -0,0 +1,161 @@
+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=Position(0), 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 not type(sel.selection) is 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 not type(sel.selection) is 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)