Files
AareDAQ/src/aare/gui/models/user_sample_model.py
T
perl_d 2c237b6623
CI / lint (push) Failing after 56s
CI / test (3.11) (push) Skipped
CI / test (3.12) (push) Skipped
CI / lint (pull_request) Failing after 2m11s
CI / test (3.11) (pull_request) Skipped
CI / test (3.12) (pull_request) Skipped
style: fix a few style issues by hand
2026-07-27 14:57:29 +02:00

329 lines
12 KiB
Python

import re
from aarecommon.models.models import SampleShortInfo, SampleShortInfoList
from PySide6.QtCore import QAbstractTableModel, QMimeData, Qt
from PySide6.QtGui import QBrush, QColor
def get_entry(sample: SampleShortInfo, column: int):
if column == 0:
return sample.sample_name
elif column == 1:
return sample.puck_name
elif column == 2:
return sample.dewar_name
elif column == 3:
return sample.loc_str()
elif column == 4:
return sample.priority
elif column == 5:
return sample.user
elif column == 6:
return sample.mount_count
elif column == 7:
return sample.raster_count
elif column == 8:
return sample.rotation_count
elif column == 9:
return sample.screening_count
elif column == 10:
return sample.comment
class UserSampleSpreadsheet(QAbstractTableModel):
def __init__(
self,
parent=None,
samples: list[SampleShortInfo] | None = None,
current_puck: str | None = None,
current_sample: int | None = None,
):
super().__init__(parent)
if samples is None:
samples = []
self.samples: list[SampleShortInfo] = samples
self.header = [
"Sample name",
"Puck",
"Dewar",
"Location",
"Priority",
"User",
"Mount count",
"Raster count",
"Rotation count",
"Screening count",
"Comment",
]
self.current_sample = current_sample
self.current_puck = current_puck
self.__sort_col = 3
self.__sort_order = Qt.SortOrder.AscendingOrder
self.__filters: dict[int, str] = {}
self.__filter_col: int | None = 5
self.__filter_value: str | None = None
self.current_pgroup: str | None = None
self.show_all_pgroups: bool = False
self._sort()
def to_list(self) -> list[dict]:
return [sample.to_dict() for sample in self.samples]
def from_list(self, data: list[dict]):
self.beginResetModel()
self.samples = [SampleShortInfo.from_dict(d) for d in data]
self.endResetModel()
def rowCount(self, parent=None):
return len(self.__sorted_samples)
def columnCount(self, parent=None):
return len(self.header)
def data(self, index, role=None):
if role == Qt.ItemDataRole.DisplayRole:
return get_entry(self.__sorted_samples[index.row()], index.column())
elif role == Qt.ItemDataRole.TextAlignmentRole: # Align text to center
return Qt.AlignmentFlag.AlignCenter
elif role == Qt.ItemDataRole.BackgroundRole:
if self.__sorted_samples[index.row()].db_id == self.current_sample:
return QBrush(QColor(114, 159, 207)) # darker blue
if self.__sorted_samples[index.row()].puck_name == self.current_puck:
return QBrush(QColor(216, 228, 253)) # light blue
return QBrush(QColor(255, 255, 255)) # White
return None # For other roles, return None
def headerData(self, section, orientation, role=None):
if role == Qt.ItemDataRole.DisplayRole:
if orientation == Qt.Orientation.Horizontal: # Column header
return self.header[section] if self.header else f"Column {section + 1}"
if orientation == Qt.Orientation.Vertical: # Row header
return str(section + 1) # Row numbers start from 1
return None
def updateCurrentSample(
self, current_puck: str | None = None, current_sample: int | None = None
):
self.current_puck = current_puck
self.current_sample = current_sample
def updateData(self, samples: list[SampleShortInfo]):
if samples != self.samples:
self.beginResetModel()
self.samples = samples
self._sort()
self.endResetModel()
def sort(self, column, order):
self.layoutAboutToBeChanged.emit()
self.__sort_order = order
self.__sort_col = column
self._sort()
self.layoutChanged.emit()
def _sort(self):
filtered = self._apply_filter(self.samples)
if self.__sort_col == 3:
self.__sorted_samples = sorted(
filtered,
key=lambda row: row.loc_str_sort(),
reverse=(
self.__sort_order == Qt.SortOrder.DescendingOrder
), # Reverse for descending order
)
else:
# Sort the samples based on the specified column and order
self.__sorted_samples = sorted(
filtered,
key=lambda row: get_entry(row, self.__sort_col),
reverse=(
self.__sort_order == Qt.SortOrder.DescendingOrder
), # Reverse for descending order
)
def _apply_filter(self, rows: list[SampleShortInfo]) -> list[SampleShortInfo]:
# Default filter by User using current p-group if no explicit filter set
filters: dict[int, str] = {
col: v for col, v in (self.__filters or {}).items() if (v or "").strip()
}
if self.__filter_col is not None and (self.__filter_value or "").strip():
filters[self.__filter_col] = self.__filter_value
if 5 not in filters and self.current_pgroup and not self.show_all_pgroups:
filters[5] = self.current_pgroup
if not filters:
return rows
def matches(sample: SampleShortInfo) -> bool:
for col, txt in filters.items():
val = get_entry(sample, col)
s = "" if val is None else str(val)
if str(txt).strip().lower() not in s.strip().lower():
return False
return True
return [r for r in rows if matches(r)]
def flags(self, index):
default_flags = super().flags(index)
if index.isValid():
return Qt.ItemFlag.ItemIsDragEnabled | default_flags
return default_flags
def mimeTypes(self):
return ["text/plain"]
def mimeData(self, indexes):
mime_data = QMimeData()
sample_data = SampleShortInfoList(s=[])
for i in sorted(set(index.row() for index in indexes)):
sample_data.s.append(self.__sorted_samples[i])
mime_data.setText(sample_data.model_dump_json())
return mime_data
def get_id(self, row: int) -> SampleShortInfo:
return self.__sorted_samples[row]
def set_filter(self, field: str, text: str | None):
try:
col = self.header.index(field)
except ValueError:
col = None
self.layoutAboutToBeChanged.emit()
self.__filter_col = col if text and text.strip() and col is not None else None
self.__filter_value = text.strip() if text else None
# mirror into multi-filter map
if col is not None:
if text and text.strip():
self.__filters[col] = text.strip()
else:
self.__filters.pop(col, None)
self._sort()
self.layoutChanged.emit()
def clear_filter(self):
self.layoutAboutToBeChanged.emit()
self.__filter_col = None
self.__filter_value = None
self.__filters.clear()
self._sort()
self.layoutChanged.emit()
def set_default_user_filter(self, pgroup: str | None):
self.layoutAboutToBeChanged.emit()
self.current_pgroup = (pgroup or "").strip() or None
self._sort()
self.layoutChanged.emit()
def set_column_filter(self, column: int, text: str | None):
self.layoutAboutToBeChanged.emit()
if text and text.strip():
self.__filters[column] = text.strip()
else:
self.__filters.pop(column, None)
self._sort()
self.layoutChanged.emit()
def clear_all_column_filters(self):
self.clear_filter()
def set_show_all_pgroups(self, show_all: bool):
self.layoutAboutToBeChanged.emit()
self.show_all_pgroups = bool(show_all)
self._sort()
self.layoutChanged.emit()
def unique_values_for_column(self, column: int, limit: int = 200) -> list[str]:
"""Get unique values for a column from currently filtered samples (excluding this column's filter)."""
# Get currently filtered samples, but exclude the filter for this column
temp_filter = self.__filters.pop(column, None)
filtered_samples = self._apply_filter(self.samples)
# Restore the filter
if temp_filter is not None:
self.__filters[column] = temp_filter
seen: set[str] = set()
out: list[str] = []
for s in filtered_samples:
v = get_entry(s, column)
if v is None:
continue
txt = str(v).strip()
if not txt:
continue
if txt not in seen:
seen.add(txt)
out.append(txt)
if len(out) >= limit:
break
# Sort appropriately
if column == 5: # User/pgroup column
try:
out.sort(
key=lambda x: (
int(x[1:])
if x and x[0].lower() == "p" and x[1:].isdigit()
else float("inf")
)
)
except Exception:
out.sort()
else:
out.sort()
return out
def suggested_prefixes_for_sample_name(self, limit: int = 200) -> list[str]:
"""Get sample name prefixes from currently filtered samples (excluding column 0 filter)."""
# Get currently filtered samples, excluding the sample name filter
temp_filter = self.__filters.pop(0, None)
filtered_samples = self._apply_filter(self.samples)
if temp_filter is not None:
self.__filters[0] = temp_filter
rx = re.compile(r"^([A-Za-z]+)")
counts: dict[str, int] = {}
for s in filtered_samples:
name = (s.sample_name or "").strip()
if not name:
continue
m = rx.match(name)
pref = m.group(1) if m else name.split("_")[0]
pref = pref.strip()
if not pref:
continue
counts[pref] = counts.get(pref, 0) + 1
# Sort by frequency desc, then alpha
items = sorted(counts.items(), key=lambda kv: (-kv[1], kv[0]))
return [k for k, _ in items[:limit]]
def suggested_prefixes_for_location(self, limit: int = 200) -> tuple[list[str], list[str]]:
"""Get location prefixes from currently filtered samples (excluding column 3 filter)."""
# Get currently filtered samples, excluding the location filter
temp_filter = self.__filters.pop(3, None)
filtered_samples = self._apply_filter(self.samples)
if temp_filter is not None:
self.__filters[3] = temp_filter
seg_seen: set[str] = set()
segpos_seen: set[str] = set()
for s in filtered_samples:
loc = s.loc_str() # "-" or like "B3-10"
if loc and loc != "-" and isinstance(loc, str) and "-" in loc:
left = loc.split("-", 1)[0] # e.g., "B3"
if left:
segpos_seen.add(left)
seg = left[0]
if seg:
seg_seen.add(seg)
segs = sorted(seg_seen)
segpos = sorted(segpos_seen, key=lambda x: (x[0], int(x[1:]) if x[1:].isdigit() else 0))
return (segs[:limit], segpos[:limit])