From c0c2f41990ec24105310c279da78082f185574ce Mon Sep 17 00:00:00 2001 From: lachlangrose Date: Wed, 30 Sep 2026 10:38:00 +0930 Subject: [PATCH 1/3] feat: add block model to the visualisation Add a Block Model button to the visualisation panel. It fills the model bounding box with blocks, evaluates the stratigraphy at each block centre on a background thread, and colours the blocks by the stratigraphic column. The number of blocks along each axis is set in a dialog and defaults to the model resolution. --- .../gui/visualisation/cross_section_utils.py | 27 +++ .../gui/visualisation/feature_list_widget.py | 178 +++++++++++++++++- .../resources/images/block_model.svg | 13 ++ tests/unit/test_block_model_mesh.py | 27 +++ 4 files changed, 244 insertions(+), 1 deletion(-) create mode 100644 loopstructural/resources/images/block_model.svg create mode 100644 tests/unit/test_block_model_mesh.py diff --git a/loopstructural/gui/visualisation/cross_section_utils.py b/loopstructural/gui/visualisation/cross_section_utils.py index 68dd5aa..43af389 100644 --- a/loopstructural/gui/visualisation/cross_section_utils.py +++ b/loopstructural/gui/visualisation/cross_section_utils.py @@ -69,3 +69,30 @@ def build_line_extrusion_mesh( yy = np.tile(xy[:, 1][:, None], (1, z_resolution)) zz = np.tile(z[None, :], (resolution, 1)) return pv.StructuredGrid(xx, yy, zz) + + +def build_block_model_mesh(origin, maximum, ncells) -> pv.ImageData: + """Build a regular voxel grid (block model) that fills a bounding box. + + Parameters + ---------- + origin, maximum : array_like + (3,) opposite corners of the box, in the model's coordinate system. + ncells : array_like + (3,) number of cells (blocks) along x, y and z. + + Returns + ------- + pv.ImageData + A grid with `prod(ncells)` cells. Callers evaluate the model at + `mesh.cell_centers().points` and store the result as cell data. + """ + origin = np.asarray(origin, dtype=float) + maximum = np.asarray(maximum, dtype=float) + ncells = np.maximum(np.asarray(ncells, dtype=int), 1) + spacing = (maximum - origin) / ncells + return pv.ImageData( + dimensions=tuple(int(n) for n in ncells + 1), + spacing=tuple(float(s) for s in spacing), + origin=tuple(float(o) for o in origin), + ) diff --git a/loopstructural/gui/visualisation/feature_list_widget.py b/loopstructural/gui/visualisation/feature_list_widget.py index b31d024..d6ea05c 100644 --- a/loopstructural/gui/visualisation/feature_list_widget.py +++ b/loopstructural/gui/visualisation/feature_list_widget.py @@ -38,7 +38,11 @@ from ..background_task import finish_background_task, start_background_task from ..compatibility import configure_layer_combo -from .cross_section_utils import build_line_extrusion_mesh, build_plane_mesh +from .cross_section_utils import ( + build_block_model_mesh, + build_line_extrusion_mesh, + build_plane_mesh, +) from .mesh_scalar_utils import stratigraphic_ids_to_rgb logger = logging.getLogger(__name__) @@ -88,6 +92,7 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager self._topography_progress = None self._build_cross_section_controls() + self._build_block_model_controls() # A single row of icon-only actions, in workflow order, replaces the # previous stack of full-width text buttons. @@ -97,6 +102,7 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager actionsRow.addWidget(self.addStratigraphicSurfacesButton) actionsRow.addWidget(self.addTopographyButton) actionsRow.addWidget(self.crossSectionButton) + actionsRow.addWidget(self.blockModelButton) actionsRow.addStretch(1) self.mainLayout.addLayout(actionsRow) self.mainLayout.addWidget(self.colourTopographyByStratigraphyCheckBox) @@ -107,6 +113,12 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager self._cross_section_worker = None self._cross_section_progress = None self._pending_cross_section_name = None + + # background task handles for the block model + self._block_model_thread = None + self._block_model_worker = None + self._block_model_progress = None + self._pending_block_model_name = None # Whether the user has hand-edited the plane's origin/normal/size -- # while False, `update_feature_list` keeps re-syncing those fields to # the model's current bounding box (see @@ -281,6 +293,79 @@ def make_coord_spinbox(default=0.0): closeButton.clicked.connect(self.crossSectionDialog.close) dialogLayout.addWidget(closeButton) + def _build_block_model_controls(self): + """Build the "Block Model" icon button (added to the shared actions + row in __init__), which opens a dialog to set the number of blocks + along each axis. The block model fills the model bounding box and is + coloured by the stratigraphic column, like the cross sections. + """ + self.blockModelButton = self._make_custom_icon_tool_button( + "block_model.svg", "Block Model..." + ) + self.blockModelButton.clicked.connect(self._show_block_model_dialog) + self._block_model_resolution_initialised = False + + self.blockModelDialog = QDialog(self) + self.blockModelDialog.setWindowTitle("Block Model") + dialogLayout = QVBoxLayout(self.blockModelDialog) + dialogLayout.addWidget(QLabel("Blocks fill the model bounding box")) + form = QFormLayout() + + def make_cells_spinbox(): + box = QSpinBox(self) + box.setRange(1, 1000) + box.setValue(50) + return box + + self.blockModelNxSpinBox = make_cells_spinbox() + self.blockModelNySpinBox = make_cells_spinbox() + self.blockModelNzSpinBox = make_cells_spinbox() + form.addRow( + "Blocks (x, y, z)", + self._hbox( + self.blockModelNxSpinBox, self.blockModelNySpinBox, self.blockModelNzSpinBox + ), + ) + self.blockModelUseModelResolutionButton = QPushButton("Use Model Resolution", self) + self.blockModelUseModelResolutionButton.setToolTip( + "Set the number of blocks to the model's interpolation grid" + ) + self.blockModelUseModelResolutionButton.clicked.connect(self._reset_block_model_resolution) + form.addRow("", self.blockModelUseModelResolutionButton) + dialogLayout.addLayout(form) + + self.addBlockModelButton = QPushButton("Add Block Model", self) + self.addBlockModelButton.clicked.connect(self.add_block_model) + dialogLayout.addWidget(self.addBlockModelButton) + + closeButton = QPushButton("Close", self.blockModelDialog) + closeButton.clicked.connect(self.blockModelDialog.close) + dialogLayout.addWidget(closeButton) + + def _show_block_model_dialog(self): + # Start from the model's own resolution the first time only, so the + # user's block counts are kept between openings of the dialog. + if not self._block_model_resolution_initialised: + self._reset_block_model_resolution() + self._block_model_resolution_initialised = True + self.blockModelDialog.show() + self.blockModelDialog.raise_() + self.blockModelDialog.activateWindow() + + def _reset_block_model_resolution(self): + if not self.model_manager or self.model_manager.model is None: + return + try: + nsteps = np.asarray(self.model_manager.model.bounding_box.nsteps, dtype=int) + except Exception: + logger.info("Model bounding box has no resolution.") + return + for box, value in zip( + (self.blockModelNxSpinBox, self.blockModelNySpinBox, self.blockModelNzSpinBox), + nsteps, + ): + box.setValue(int(value)) + def _show_cross_section_dialog(self): self.crossSectionDialog.show() self.crossSectionDialog.raise_() @@ -950,6 +1035,97 @@ def _on_cross_section_error(self, traceback_text): self.addLineCrossSectionButton.setEnabled(True) logger.error(f"Failed to build cross section: {traceback_text}") + def add_block_model(self): + """Fill the model bounding box with blocks and colour each block by + the stratigraphic unit at its centre. + + Evaluating the model at every block centre can be slow, so this runs + on a background thread (see `add_topography_surface`). + """ + if not self.model_manager: + logger.info("Model manager is not set.") + return + if self.model_manager.model is None: + logger.info("No model available to build a block model.") + return + missing = self.model_manager.get_units_without_colour() + if missing: + QMessageBox.warning( + self, + "Missing unit colour", + "Cannot colour the block model by stratigraphy. These units have no " + "valid colour in the stratigraphic column:\n\n" + + "\n".join(missing) + + "\n\nSet a colour for each unit and try again.", + ) + return + + bb = self.model_manager.model.bounding_box + origin = np.asarray(bb.origin, dtype=float) + maximum = np.asarray(bb.maximum, dtype=float) + ncells = ( + self.blockModelNxSpinBox.value(), + self.blockModelNySpinBox.value(), + self.blockModelNzSpinBox.value(), + ) + + def target(progress_callback): + progress_callback("Building block model grid...") + mesh = build_block_model_mesh(origin, maximum, ncells) + progress_callback("Evaluating stratigraphy on block model...") + ids = self.model_manager.evaluate_stratigraphy_on_points(mesh.cell_centers().points) + colours = self.model_manager.get_stratigraphic_column_colours() + return mesh, ids, colours + + self._pending_block_model_name = self._unique_cross_section_name('block_model') + self.addBlockModelButton.setEnabled(False) + self._block_model_thread, self._block_model_worker, self._block_model_progress = ( + start_background_task( + self, + target, + title="Block Model", + initial_label="Building block model grid...", + on_progress=self._on_block_model_progress, + on_finished=self._on_block_model_finished, + on_error=self._on_block_model_error, + ) + ) + + def _on_block_model_progress(self, message): + try: + self._block_model_progress.setLabelText(message) + except Exception: + pass + + def _on_block_model_finished(self, result): + finish_background_task( + self._block_model_thread, self._block_model_worker, self._block_model_progress + ) + self.addBlockModelButton.setEnabled(True) + + mesh, ids, colours = result + # Keep the unit ids on the mesh so the object properties panel can + # also colour or threshold the blocks by unit. + mesh.cell_data['stratigraphy'] = np.asarray(ids) + mesh.cell_data['colour'] = stratigraphic_ids_to_rgb(ids, colours) + self.viewer.add_mesh_object( + mesh, + name=self._pending_block_model_name, + scalars='colour', + rgb=True, + show_scalar_bar=False, + show_edges=False, + source_type='block_model', + ) + logger.info(f"Added block model '{self._pending_block_model_name}'.") + + def _on_block_model_error(self, traceback_text): + finish_background_task( + self._block_model_thread, self._block_model_worker, self._block_model_progress + ) + self.addBlockModelButton.setEnabled(True) + logger.error(f"Failed to build block model: {traceback_text}") + def _on_model_update(self, event: str, *args): """Called when the underlying model_manager notifies observers. diff --git a/loopstructural/resources/images/block_model.svg b/loopstructural/resources/images/block_model.svg new file mode 100644 index 0000000..14dfd21 --- /dev/null +++ b/loopstructural/resources/images/block_model.svg @@ -0,0 +1,13 @@ + + + + + + + + + + + + + diff --git a/tests/unit/test_block_model_mesh.py b/tests/unit/test_block_model_mesh.py new file mode 100644 index 0000000..e7ea9a6 --- /dev/null +++ b/tests/unit/test_block_model_mesh.py @@ -0,0 +1,27 @@ +"""Pytest tests for `build_block_model_mesh` (see +loopstructural/gui/visualisation/cross_section_utils.py). + +Only pyvista/numpy are needed, so these run in the fast tests/unit/ job. +""" + +import numpy as np + +from loopstructural.gui.visualisation.cross_section_utils import build_block_model_mesh + + +def test_block_model_fills_bounding_box(): + mesh = build_block_model_mesh((0, 10, -100), (100, 60, 0), (10, 5, 4)) + assert mesh.n_cells == 10 * 5 * 4 + np.testing.assert_allclose(mesh.bounds, (0, 100, 10, 60, -100, 0)) + + +def test_block_model_cell_centres_are_inside_box(): + mesh = build_block_model_mesh((0, 0, 0), (10, 10, 10), (2, 2, 2)) + centres = mesh.cell_centers().points + assert centres.shape == (8, 3) + np.testing.assert_allclose(np.unique(centres[:, 0]), [2.5, 7.5]) + + +def test_block_model_needs_at_least_one_cell_per_axis(): + mesh = build_block_model_mesh((0, 0, 0), (1, 1, 1), (0, 3, 3)) + assert mesh.n_cells == 1 * 3 * 3 From 397cad8bccb4e7169a4dd969c836a6509d922990 Mon Sep 17 00:00:00 2001 From: lachlangrose Date: Wed, 30 Sep 2026 12:05:02 +0930 Subject: [PATCH 2/3] feat: update out-of-date viewer objects with a button Model changes no longer rebuild the viewer objects automatically. The objects built from the model are marked out of date (grey italic in the object list), and a new button builds them again from the current model on a background thread. Each object keeps its name, visibility and viewer settings. Cross sections, block models and topography are also updated. Also keep the source information when the properties panel adds an object again, and save opacity changes. --- .../gui/visualisation/feature_list_widget.py | 675 +++++++++++------- .../visualisation/loop_pyvistaqt_wrapper.py | 45 ++ .../gui/visualisation/object_list_widget.py | 12 +- .../visualisation/object_properties_widget.py | 17 +- 4 files changed, 467 insertions(+), 282 deletions(-) diff --git a/loopstructural/gui/visualisation/feature_list_widget.py b/loopstructural/gui/visualisation/feature_list_widget.py index d6ea05c..303ee4c 100644 --- a/loopstructural/gui/visualisation/feature_list_widget.py +++ b/loopstructural/gui/visualisation/feature_list_widget.py @@ -107,6 +107,20 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager self.mainLayout.addLayout(actionsRow) self.mainLayout.addWidget(self.colourTopographyByStratigraphyCheckBox) + # Objects in the viewer are not rebuilt automatically when the model + # changes (that can re-solve the whole model after each small edit). + # They are marked out of date, and this button rebuilds them. + self.updateObjectsButton = QPushButton(self) + self.updateObjectsButton.setIcon(QgsApplication.getThemeIcon("mActionRefresh.svg")) + self.updateObjectsButton.clicked.connect(self.update_out_of_date_objects) + self.mainLayout.addWidget(self.updateObjectsButton) + self._update_objects_thread = None + self._update_objects_worker = None + self._update_objects_progress = None + if self.viewer is not None: + self.viewer.outOfDateChanged.connect(self._refresh_update_objects_button) + self._refresh_update_objects_button() + # background task handles shared by the plane and line cross-section # actions (only one can run at a time) self._cross_section_thread = None @@ -140,8 +154,8 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager self._disp_update = self.model_manager.attach( lambda _obs, _event, *a, **k: self.update_feature_list(), 'model_updated' ) - # also listen for model and feature updates so visualisation can refresh - # forward event and args into the handler so it can act on specific surfaces + # also listen for model and feature updates so the viewer + # objects built from the model can be marked out of date self._disp_feature = self.model_manager.attach( lambda _obs, _event, *a, **k: self._on_model_update(_event, *a), 'model_updated' ) @@ -458,10 +472,12 @@ def contextMenuEvent(self, event): elif action == add_data_action: self.add_data(feature_name) + def _build_scalar_field(self, feature_name): + return self.model_manager.model[feature_name].scalar_field().vtk() + def add_scalar_field(self, feature_name): - scalar_field = self.model_manager.model[feature_name].scalar_field() self.viewer.add_mesh_object( - scalar_field.vtk(), + self._build_scalar_field(feature_name), name=f'{feature_name}_scalar_field', source_feature=feature_name, source_type='feature_scalar', @@ -495,11 +511,20 @@ def add_surface(self, feature_name): isovalue=isovalue, ) - def add_vector_field(self, feature_name): + def _build_feature_surface(self, feature_name, isovalue): + feature = self.model_manager.model[feature_name] + surfaces = feature.surfaces(isovalue) if isovalue is not None else feature.surfaces() + if not surfaces: + raise ValueError(f"Feature '{feature_name}' has no surface at {isovalue}") + return surfaces[0].vtk() + + def _build_vector_field(self, feature_name): vector_field = self.model_manager.model[feature_name].vector_field() - scale = self._get_vector_scale() + return vector_field.vtk(scale=self._get_vector_scale()) + + def add_vector_field(self, feature_name): self.viewer.add_mesh_object( - vector_field.vtk(scale=scale), + self._build_vector_field(feature_name), name=f'{feature_name}_vector_field', source_feature=feature_name, source_type='feature_vector', @@ -530,23 +555,18 @@ def _get_fold(self, feature_name): fold = getattr(getattr(feature, 'builder', None), 'fold', None) return fold - def add_fold_constraint(self, feature_name, constraint): - """Add the vectors of one fold constraint of a folded feature to the viewer. + def _build_fold_constraint(self, feature_name, constraint): + """Return the vectors of one fold constraint of a folded feature as + a mesh, or None if the feature is not folded or the vectors are not + defined. The vectors are evaluated on the model grid, in the same way as the interpolator evaluates them on the element barycentres. - - Parameters - ---------- - feature_name : str - Name of the folded feature. - constraint : str - One of 'direction', 'axis' or 'norm'. """ fold = self._get_fold(feature_name) if fold is None: logger.info(f"Feature {feature_name} is not folded") - return + return None feature = self.model_manager.model[feature_name] # make sure the fold rotation angles are fitted feature.builder.up_to_date() @@ -559,53 +579,75 @@ def add_fold_constraint(self, feature_name, constraint): logger.warning( f"Fold {constraint} vectors do not match the grid points ({vectors.shape} != {points.shape})" ) - return + return None length = np.linalg.norm(vectors, axis=1) mask = np.all(np.isfinite(vectors), axis=1) & (length > 0) if not np.any(mask): logger.warning(f"Fold {constraint} vectors for {feature_name} are not defined") - return + return None vectors = vectors[mask] / length[mask, None] vector_points = VectorPoints(points[mask], vectors, f'{feature_name}_fold_{constraint}') + return vector_points.vtk(scale=self._get_vector_scale()) + + def add_fold_constraint(self, feature_name, constraint): + """Add the vectors of one fold constraint of a folded feature to the viewer. + + Parameters + ---------- + feature_name : str + Name of the folded feature. + constraint : str + One of 'direction', 'axis' or 'norm'. + """ + mesh = self._build_fold_constraint(feature_name, constraint) + if mesh is None: + return self.viewer.add_mesh_object( - vector_points.vtk(scale=self._get_vector_scale()), + mesh, name=f'{feature_name}_fold_{constraint}', color=self.FOLD_CONSTRAINT_COLOURS[constraint], source_feature=feature_name, source_type=f'fold_constraint_{constraint}', ) - def add_data(self, feature_name): - data = self.model_manager.model[feature_name].get_data() - for d in data: + def _build_data_meshes(self, feature_name): + """Return (name, mesh, source_type) for each data set of a feature.""" + meshes = [] + for d in self.model_manager.model[feature_name].get_data(): d.locations = self.model_manager.model.rescale(d.locations) if issubclass(type(d), VectorPoints): - scale = self._get_vector_scale() # tolerance is None means all points are shown - self.viewer.add_mesh_object( - d.vtk(scale=scale, tolerance=None), - name=f'{feature_name}_{d.name}_points', - source_feature=feature_name, - source_type='feature_points', + meshes.append( + ( + f'{feature_name}_{d.name}_points', + d.vtk(scale=self._get_vector_scale(), tolerance=None), + 'feature_points', + ) ) else: - self.viewer.add_mesh_object( - d.vtk(), - name=f'{feature_name}_{d.name}', - source_feature=feature_name, - source_type='feature_data', - ) + meshes.append((f'{feature_name}_{d.name}', d.vtk(), 'feature_data')) + return meshes + + def add_data(self, feature_name): + for name, mesh, source_type in self._build_data_meshes(feature_name): + self.viewer.add_mesh_object( + mesh, name=name, source_feature=feature_name, source_type=source_type + ) logger.info(f"Adding data to feature: {feature_name}") + def _build_bounding_box(self): + return self.model_manager.model.bounding_box.vtk().outline() + def add_model_bounding_box(self): if not self.model_manager: logger.info("Model manager is not set.") return - bb = self.model_manager.model.bounding_box.vtk().outline() self.viewer.add_mesh_object( - bb, name='model_bounding_box', source_feature='__model__', source_type='bounding_box' + self._build_bounding_box(), + name='model_bounding_box', + source_feature='__model__', + source_type='bounding_box', ) - # Logic for adding model bounding box logger.info("Adding model bounding box...") def add_fault_surfaces(self): @@ -624,6 +666,18 @@ def add_fault_surfaces(self): ) logger.info("Adding fault surfaces...") + def _find_fault_surface(self, name): + for surface in self.model_manager.model.get_fault_surfaces(): + if str(surface.name) == str(name): + return surface + raise ValueError(f"Fault surface '{name}' is not in the model") + + def _find_stratigraphic_surface(self, name): + for surface in self.model_manager.model.get_stratigraphic_surfaces(): + if str(surface.name) == str(name): + return surface + raise ValueError(f"Stratigraphic surface '{name}' is not in the model") + def add_stratigraphic_surfaces(self): if not self.model_manager: logger.info("Model manager is not set.") @@ -704,6 +758,7 @@ def _on_topography_grid_finished(self, result): cmap='terrain', show_scalar_bar=True, source_type='topography_surface', + metadata={'coloured': False}, ) self.colourTopographyByStratigraphyCheckBox.setEnabled(True) if self.colourTopographyByStratigraphyCheckBox.isChecked(): @@ -734,6 +789,7 @@ def _on_colour_topography_toggled(self, checked): cmap='terrain', show_scalar_bar=True, source_type='topography_surface', + metadata={'coloured': False}, ) def _colour_topography_surface(self): @@ -789,11 +845,11 @@ def _on_topography_colour_finished(self, result): ids, colours = result mesh = self.viewer.meshes['topography_surface']['mesh'] - rgb = stratigraphic_ids_to_rgb(ids, colours) + self._set_stratigraphy_arrays(mesh.point_data, ids, colours) self.viewer.add_mesh_object( mesh, name='topography_surface', - scalars=rgb, + scalars='colour', rgb=True, show_scalar_bar=False, # Directional lighting shades the same colour differently depending @@ -803,6 +859,7 @@ def _on_topography_colour_finished(self, result): # keeps it a true, direct colour-for-colour match. lighting=False, source_type='topography_surface', + metadata={'coloured': True}, ) logger.info("Coloured topography surface by stratigraphic column.") @@ -998,11 +1055,11 @@ def _on_cross_section_progress(self, message): def _add_cross_section_mesh(self, result, source_type): mesh, ids, colours = result - rgb = stratigraphic_ids_to_rgb(ids, colours) + self._set_stratigraphy_arrays(mesh.point_data, ids, colours) self.viewer.add_mesh_object( mesh, name=self._pending_cross_section_name, - scalars=rgb, + scalars='colour', rgb=True, show_scalar_bar=False, # see `_on_topography_colour_finished` for why cross sections use @@ -1104,10 +1161,7 @@ def _on_block_model_finished(self, result): self.addBlockModelButton.setEnabled(True) mesh, ids, colours = result - # Keep the unit ids on the mesh so the object properties panel can - # also colour or threshold the blocks by unit. - mesh.cell_data['stratigraphy'] = np.asarray(ids) - mesh.cell_data['colour'] = stratigraphic_ids_to_rgb(ids, colours) + self._set_stratigraphy_arrays(mesh.cell_data, ids, colours) self.viewer.add_mesh_object( mesh, name=self._pending_block_model_name, @@ -1116,6 +1170,7 @@ def _on_block_model_finished(self, result): show_scalar_bar=False, show_edges=False, source_type='block_model', + metadata={'ncells': [int(n) for n in np.asarray(mesh.dimensions) - 1]}, ) logger.info(f"Added block model '{self._pending_block_model_name}'.") @@ -1126,244 +1181,320 @@ def _on_block_model_error(self, traceback_text): self.addBlockModelButton.setEnabled(True) logger.error(f"Failed to build block model: {traceback_text}") - def _on_model_update(self, event: str, *args): - """Called when the underlying model_manager notifies observers. - - We remove any meshes that were created from model features and re-add - them from the current model so visualisation follows model changes. - - If the notification is for a specific feature (event == 'feature_updated') - and an isovalue is provided (either as second arg or stored in viewer - metadata), only the matching surface will be re-added. For generic - 'model_updated' notifications the previous behaviour (re-add all - affected feature representations) is preserved. + @staticmethod + def _set_stratigraphy_arrays(data, ids, colours): + """Store the unit ids and their colours as the 'stratigraphy' and + 'colour' arrays of `data` (a mesh's point_data or cell_data). + + Named arrays (not an RGB array passed straight to the viewer) let the + object properties panel use the unit ids, and let + `update_out_of_date_objects` add the object again with the same + viewer settings. """ + data['stratigraphy'] = np.asarray(ids) + data['colour'] = stratigraphic_ids_to_rgb(ids, colours) + + def _colour_by_stratigraphy(self, points, data): + ids = self.model_manager.evaluate_stratigraphy_on_points(points) + colours = self.model_manager.get_stratigraphic_column_colours() + self._set_stratigraphy_arrays(data, ids, colours) + + # Viewer objects that `_rebuild_object` can build again from the model, + # by `source_type` (fold constraints use the 'fold_constraint_' prefix). + REBUILDABLE_SOURCE_TYPES = { + 'feature_scalar', + 'feature_surface', + 'feature_vector', + 'feature_vectors', + 'feature_points', + 'feature_data', + 'bounding_box', + 'fault_surface', + 'stratigraphic_surface', + 'cross_section_plane', + 'cross_section_line', + 'block_model', + 'topography_surface', + } + # Source types that are coloured by the stratigraphic column + STRATIGRAPHY_COLOURED_SOURCE_TYPES = { + 'cross_section_plane', + 'cross_section_line', + 'block_model', + } + # Source types that are built from one feature of the model + FEATURE_SOURCE_TYPES = { + 'feature_scalar', + 'feature_surface', + 'feature_vector', + 'feature_vectors', + 'feature_points', + 'feature_data', + } - # Prefer the DebugManager for logging when available (it forwards to - # the plugin/toolbelt logger and handles debug mode). Fall back to the - # module logger if no debug manager is present. - def _log(msg, level=0): - try: - dbg = None - if getattr(self, 'model_manager', None) is not None: - dbg = getattr(self.model_manager, '_debug_manager', None) - if dbg is not None and hasattr(dbg, 'log'): - # DebugManager.log expects message and log_level keyword - dbg.log(str(msg), log_level=level) - else: - logger.info(str(msg)) - except Exception: - try: - logger.info(str(msg)) - except Exception: - pass + def _is_rebuildable(self, meta) -> bool: + source_type = meta.get('source_type') or '' + return source_type in self.REBUILDABLE_SOURCE_TYPES or source_type.startswith( + 'fold_constraint_' + ) - _log(f"Model update event received: {event} with args: {args}") - try: - _log([f"Mesh: {name}, Meta: {meta}" for name, meta in self.viewer.meshes.items()]) - except Exception: - _log("Model update: failed to enumerate viewer meshes") + def _uses_stratigraphy_colours(self, spec) -> bool: + if spec['source_type'] in self.STRATIGRAPHY_COLOURED_SOURCE_TYPES: + return True + return spec['source_type'] == 'topography_surface' and bool( + spec['metadata'].get('coloured') + ) + + def _on_model_update(self, event: str, *args): + """Mark the viewer objects built from the model as out of date. - if not self.model_manager or not self.viewer: + The objects are not rebuilt here: rebuilding can solve the model + again, which is slow after each small edit. The user rebuilds them + with the update button (see `update_out_of_date_objects`). + """ + if not self.viewer: return if event not in ('model_updated', 'feature_updated'): return - feature_name = None - if event == 'feature_updated' and len(args) >= 1: - feature_name = args[0] + names = [ + name for name, meta in list(self.viewer.meshes.items()) if self._is_rebuildable(meta) + ] + self.viewer.set_out_of_date(names, True) + + def _refresh_update_objects_button(self): + count = len(self.viewer.out_of_date_objects()) if self.viewer is not None else 0 + if self._update_objects_thread is not None: + self.updateObjectsButton.setText("Updating Viewer Objects...") + self.updateObjectsButton.setEnabled(False) + elif count: + noun = "Object" if count == 1 else "Objects" + self.updateObjectsButton.setText(f"Update {count} Out-of-Date {noun}") + self.updateObjectsButton.setToolTip( + "The model changed after these objects were added to the viewer. " + "Build them again from the current model." + ) + self.updateObjectsButton.setEnabled(True) + else: + self.updateObjectsButton.setText("Viewer Objects Up to Date") + self.updateObjectsButton.setToolTip("") + self.updateObjectsButton.setEnabled(False) - # If the model was reset (None) or features referenced by viewer meshes - # no longer exist in the current model, remove the linkage from those - # meshes so they are not treated as feature-driven on subsequent updates. - try: - try: - current_features = {f.name for f in self.model_manager.features()} - except Exception: - current_features = set() - - # If the model is None or a feature referenced by a mesh is missing, - # decouple that mesh from the feature so it remains visible but won't - # be auto-updated or re-added when the model changes. - for mesh_name, meta in list(self.viewer.meshes.items()): - sf = meta.get('source_feature', None) - if sf is None: - continue - if self.model_manager.model is None or sf not in current_features: - _log(f"Decoupling mesh '{mesh_name}' from missing feature '{sf}'") - meta.pop('source_feature', None) - meta.pop('source_type', None) - meta.pop('isovalue', None) - # mark as decoupled so other logic can detect it if needed - meta['decoupled_from_feature'] = True - except Exception: - _log('Failed while decoupling meshes from features') - - # Build a set of features that currently have viewer meshes - affected_features = set() - for _, meta in list(self.viewer.meshes.items()): - if feature_name is not None: - if meta.get('source_feature', None) == feature_name: - affected_features.add(feature_name) - _log(f"Updating visualisation for feature: {feature_name}") - continue - - sf = meta.get('source_feature', None) - - if sf is not None: - affected_features.add(sf) - _log(f"Affected features to update: {affected_features}") - # For each affected feature, only update existing meshes tied to that feature - for feature_name in affected_features: - # collect mesh names that belong to this feature (snapshot to avoid mutation while iterating) - meshes_for_feature = [ - name - for name, meta in list(self.viewer.meshes.items()) - if meta.get('source_feature') == feature_name - ] - _log(f"Re-adding meshes for feature: {feature_name}: {meshes_for_feature}") - - for mesh_name in meshes_for_feature: - meta = self.viewer.meshes.get(mesh_name, {}) - source_type = meta.get('source_type') - kwargs = meta.get('kwargs', {}) or {} - isovalue = meta.get('isovalue', None) - - # remove existing actor/entry so add_mesh_object can recreate with same name - try: - self.viewer.remove_object(mesh_name) - _log(f"Removed existing mesh: {mesh_name}") - except Exception: - _log(f"Failed to remove existing mesh: {mesh_name}") + def update_out_of_date_objects(self): + """Build all out-of-date viewer objects again from the current model. - try: - # Surfaces associated with individual features - if source_type == 'feature_surface': - surfaces = [] - try: - if isovalue is not None: - surfaces = self.model_manager.model[feature_name].surfaces(isovalue) - else: - surfaces = self.model_manager.model[feature_name].surfaces() - - if surfaces: - add_name = mesh_name - _log( - f"Re-adding surface for feature: {feature_name} with isovalue: {isovalue} and {kwargs}" - ) - kwargs['isovalue'] = isovalue - - self.viewer.add_mesh_object( - surfaces[0].vtk(), - name=add_name, - source_feature=feature_name, - source_type='feature_surface', - isovalue=isovalue, - **kwargs, - ) - continue - except Exception as e: - _log( - f"Failed to find matching surface for feature: {feature_name} with isovalue: {isovalue}, trying all surfaces. Error: {e}" - ) - - # Fault surfaces (added via add_fault_surfaces) - if source_type == 'fault_surface': - try: - fault_surfaces = self.model_manager.model.get_fault_surfaces() - match = next( - (s for s in fault_surfaces if str(s.name) == str(feature_name)), - None, - ) - if match is not None: - _log(f"Re-adding fault surface for: {feature_name}") - self.viewer.add_mesh_object( - match.vtk(), - name=mesh_name, - source_feature=feature_name, - source_type='fault_surface', - isovalue=meta.get('isovalue', 0.0), - **kwargs, - ) - continue - except Exception as e: - _log(f"Failed to re-add fault surface for {feature_name}: {e}") - - # Stratigraphic surfaces (added via add_stratigraphic_surfaces) - if source_type == 'stratigraphic_surface': - try: - strat_surfaces = self.model_manager.model.get_stratigraphic_surfaces() - match = next( - (s for s in strat_surfaces if str(s.name) == str(feature_name)), - None, - ) - if match is not None: - _log(f"Re-adding stratigraphic surface for: {feature_name}") - kwargs['color'] = getattr(match, 'colour', None) - - self.viewer.add_mesh_object( - match.vtk(), - name=mesh_name, - source_feature=feature_name, - source_type='stratigraphic_surface', - **kwargs, - ) - continue - except Exception as e: - _log(f"Failed to re-add stratigraphic surface for {feature_name}: {e}") - - # Vectors, points, scalar fields and other feature related objects - if source_type == 'feature_vector' or source_type == 'feature_vectors': - try: - self.add_vector_field(feature_name) - continue - except Exception as e: - _log(f"Failed to re-add vector field for {feature_name}: {e}") - - if source_type and source_type.startswith('fold_constraint_'): - try: - self.add_fold_constraint( - feature_name, source_type[len('fold_constraint_') :] - ) - continue - except Exception as e: - _log(f"Failed to re-add fold constraint for {feature_name}: {e}") - - if source_type in ('feature_points', 'feature_data'): - try: - self.add_data(feature_name) - continue - except Exception as e: - _log(f"Failed to re-add data for {feature_name}: {e}") - - if source_type == 'feature_scalar': - try: - self.add_scalar_field(feature_name) - continue - except Exception as e: - _log(f"Failed to re-add scalar field for {feature_name}: {e}") - - if source_type == 'bounding_box' or mesh_name == 'model_bounding_box': - try: - self.add_model_bounding_box() - continue - except Exception as e: - _log(f"Failed to re-add bounding box: {e}") - - # Fallback: if nothing matched, attempt to re-add by using viewer metadata - # Many viewer entries store the vtk source under meta['vtk'] or similar; try best-effort - try: - vtk_src = meta.get('vtk') - if vtk_src is not None: - _log(f"Fallback re-add for mesh {mesh_name}") - self.viewer.add_mesh_object(vtk_src, name=mesh_name, **kwargs) - except Exception: - pass + The objects are built on a background thread (this can solve the + model), then added to the viewer again on the GUI thread with the + same name and viewer settings (see `_on_update_objects_finished`). + """ + if not self.model_manager or self.viewer is None: + return + if self.model_manager.model is None: + logger.info("No model available to update the viewer objects.") + return + names = self.viewer.out_of_date_objects() + if not names: + return + specs = [] + for name in names: + meta = self.viewer.meshes[name] + spec = { + 'name': name, + 'source_type': meta.get('source_type') or '', + 'source_feature': meta.get('source_feature'), + 'isovalue': meta.get('isovalue'), + 'metadata': dict(meta.get('metadata') or {}), + } + if spec['source_type'] in ('cross_section_plane', 'cross_section_line'): + # the section geometry does not change; copy it so the + # background thread does not change the mesh on screen + spec['mesh'] = meta['mesh'].copy() + specs.append(spec) + + if any(self._uses_stratigraphy_colours(spec) for spec in specs): + missing = self.model_manager.get_units_without_colour() + if missing: + QMessageBox.warning( + self, + "Missing unit colour", + "Cannot update the objects coloured by stratigraphy. These units " + "have no valid colour in the stratigraphic column:\n\n" + + "\n".join(missing) + + "\n\nSet a colour for each unit and try again.", + ) + return + def target(progress_callback): + results = [] + for i, spec in enumerate(specs): + progress_callback(f"Updating {spec['name']} ({i + 1} of {len(specs)})...") + try: + mesh, overrides = self._rebuild_object(spec) + results.append((spec['name'], mesh, overrides, None)) except Exception as e: - _log(f"Failed to update visualisation for feature: {feature_name}. Error: {e}") + results.append((spec['name'], None, {}, str(e))) + return results + + self._update_objects_thread, self._update_objects_worker, self._update_objects_progress = ( + start_background_task( + self, + target, + title="Update Viewer Objects", + initial_label="Updating viewer objects...", + on_progress=self._on_update_objects_progress, + on_finished=self._on_update_objects_finished, + on_error=self._on_update_objects_error, + ) + ) + self._refresh_update_objects_button() + + def _rebuild_object(self, spec): + """Build one viewer object again from the current model. + + Runs on a background thread, so it must not touch the viewer. + Returns (mesh, overrides), where overrides are viewer settings that + come from the model (e.g. a unit colour). Raises if the object cannot + be built. + """ + source_type = spec['source_type'] + feature_name = spec['source_feature'] + metadata = spec['metadata'] + model = self.model_manager.model + + if source_type in self.FEATURE_SOURCE_TYPES or source_type.startswith('fold_constraint_'): + if feature_name is None or model.get_feature_by_name(feature_name) is None: + raise ValueError(f"Feature '{feature_name}' is not in the model") + + overrides = {} + if source_type == 'feature_scalar': + mesh = self._build_scalar_field(feature_name) + elif source_type == 'feature_surface': + mesh = self._build_feature_surface(feature_name, spec['isovalue']) + elif source_type in ('feature_vector', 'feature_vectors'): + mesh = self._build_vector_field(feature_name) + elif source_type.startswith('fold_constraint_'): + constraint = source_type[len('fold_constraint_') :] + mesh = self._build_fold_constraint(feature_name, constraint) + if mesh is None: + raise ValueError("The fold constraint vectors are not defined") + elif source_type in ('feature_points', 'feature_data'): + meshes = {name: m for name, m, _ in self._build_data_meshes(feature_name)} + if spec['name'] not in meshes: + raise ValueError(f"Feature '{feature_name}' has no data for this object") + mesh = meshes[spec['name']] + elif source_type == 'bounding_box': + mesh = self._build_bounding_box() + elif source_type == 'fault_surface': + mesh = self._find_fault_surface(feature_name).vtk() + elif source_type == 'stratigraphic_surface': + surface = self._find_stratigraphic_surface(feature_name) + mesh = surface.vtk() + overrides['color'] = surface.colour + elif source_type in ('cross_section_plane', 'cross_section_line'): + mesh = spec['mesh'] + self._colour_by_stratigraphy(mesh.points, mesh.point_data) + elif source_type == 'block_model': + bb = model.bounding_box + mesh = build_block_model_mesh(bb.origin, bb.maximum, metadata['ncells']) + self._colour_by_stratigraphy(mesh.cell_centers().points, mesh.cell_data) + elif source_type == 'topography_surface': + xx, yy, zz = self.model_manager.sample_dem_grid() + mesh = pv.StructuredGrid(xx, yy, zz) + mesh['Elevation'] = mesh.points[:, 2] + if metadata.get('coloured'): + self._colour_by_stratigraphy(mesh.points, mesh.point_data) + else: + raise ValueError(f"Cannot update objects of type '{source_type}'") + + if getattr(mesh, 'n_points', 1) == 0: + raise ValueError("The object has no geometry in the current model") + return mesh, overrides + + def _on_update_objects_progress(self, message): + try: + self._update_objects_progress.setLabelText(message) + except Exception: + pass - # Refresh the viewer + def _finish_update_objects_task(self): + finish_background_task( + self._update_objects_thread, self._update_objects_worker, self._update_objects_progress + ) + self._update_objects_thread = None + self._update_objects_worker = None + self._update_objects_progress = None + + def _on_update_objects_finished(self, results): + self._finish_update_objects_task() + failed = [] + for name, mesh, overrides, error in results: + entry = self.viewer.meshes.get(name) + if entry is None: + # removed from the viewer while the update ran + continue + if error is not None: + failed.append(f"{name}: {error}") + continue + if not self._replace_viewer_object(name, entry, mesh, overrides): + failed.append(f"{name}: cannot add the new object to the viewer") try: - self.viewer.update() + self.viewer.render() except Exception: pass + self._refresh_update_objects_button() + if failed: + logger.warning("Cannot update viewer objects:\n" + "\n".join(failed)) + QMessageBox.warning( + self, + "Update Viewer Objects", + "These objects were not updated and are still out of date:\n\n" + "\n".join(failed), + ) + + def _replace_viewer_object(self, name, entry, mesh, overrides) -> bool: + """Put `mesh` in the viewer in place of the object `name`, with the + same source values, viewer settings and visibility.""" + source = self.viewer.get_source_metadata(name) + source['out_of_date'] = False + kwargs = { + key: value + for key, value in (entry.get('kwargs') or {}).items() + if key not in source and key != 'name' + } + kwargs.update(overrides) + # a colour picked in the object properties panel has priority + user_colour = entry.get('color') + if user_colour is not None: + kwargs['color'] = user_colour + actor = entry.get('actor') + visible = bool(getattr(actor, 'visibility', True)) + + # pyvista replaces the actor with the same name, so the old object + # stays in the viewer if the new one cannot be added + try: + self.viewer.add_mesh_object(mesh, name=name, **source, **kwargs) + except Exception: + # e.g. a scalar array selected in the properties panel that the + # new mesh does not have; add it with the default colouring + for key in ('scalars', 'cmap', 'clim', 'rgb'): + kwargs.pop(key, None) + try: + self.viewer.add_mesh_object(mesh, name=name, **source, **kwargs) + except Exception: + logger.exception(f"Cannot add updated object '{name}' to the viewer") + return False + + new_entry = self.viewer.meshes.get(name, {}) + if user_colour is not None: + new_entry['color'] = user_colour + if not visible and new_entry.get('actor') is not None: + new_entry['actor'].visibility = False + return True + + def _on_update_objects_error(self, traceback_text): + self._finish_update_objects_task() + self._refresh_update_objects_button() + logger.error(f"Failed to update viewer objects: {traceback_text}") + QMessageBox.warning( + self, + "Update Viewer Objects", + "Cannot update the viewer objects. See the log for details.", + ) diff --git a/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py b/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py index b4d338c..97fab5d 100644 --- a/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py +++ b/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py @@ -6,6 +6,8 @@ class LoopPyVistaQTPlotter(QtInteractor): objectAdded = pyqtSignal(QtInteractor) # Signal to request deletion + # emitted when objects are marked out of date, or brought up to date + outOfDateChanged = pyqtSignal() def __init__(self, parent): super().__init__(parent=parent) @@ -42,6 +44,8 @@ def add_mesh_object( source_feature: Optional[str] = None, source_type: Optional[str] = None, isovalue: Optional[float] = None, + metadata: Optional[Dict[str, Any]] = None, + out_of_date: bool = False, **kwargs, ) -> None: """Add a mesh to the plotter. @@ -73,6 +77,11 @@ def add_mesh_object( source_type : Optional[str] A short tag describing the kind of source (e.g. 'feature_surface', 'fault_surface', 'bounding_box'). + metadata : Optional[dict] + Extra values needed to build the mesh again from the model (for + example the number of blocks of a block model). + out_of_date : bool + True if the mesh no longer matches the model. Returns ------- @@ -125,9 +134,43 @@ def add_mesh_object( 'source_feature': source_feature, 'source_type': source_type, 'isovalue': isovalue, + 'metadata': dict(metadata or {}), + 'out_of_date': out_of_date, } self.objectAdded.emit(self) + def get_source_metadata(self, name: str) -> Dict[str, Any]: + """Return the source values of an object as keyword arguments for + `add_mesh_object`, so that an object removed and added again (for + example to change its colour map) can still be built again from the + model. + """ + entry = self.meshes.get(name) + if not entry: + return {} + return { + 'source_feature': entry.get('source_feature'), + 'source_type': entry.get('source_type'), + 'isovalue': entry.get('isovalue'), + 'metadata': entry.get('metadata'), + 'out_of_date': bool(entry.get('out_of_date', False)), + } + + def set_out_of_date(self, names, out_of_date: bool = True) -> None: + """Mark the named objects as out of date (or up to date).""" + changed = False + for name in names: + entry = self.meshes.get(name) + if entry is not None and bool(entry.get('out_of_date')) != out_of_date: + entry['out_of_date'] = out_of_date + changed = True + if changed: + self.outOfDateChanged.emit() + + def out_of_date_objects(self): + """Return the names of the objects that are out of date.""" + return [name for name, entry in self.meshes.items() if entry.get('out_of_date')] + def remove_object(self, name: str) -> None: """Remove an object by name and clean up stored metadata. @@ -157,6 +200,8 @@ def remove_object(self, name: str) -> None: del self.meshes[name] except Exception: pass + if entry.get('out_of_date'): + self.outOfDateChanged.emit() def set_object_visibility(self, name: str, visibility): """Change the visibility of an object.""" diff --git a/loopstructural/gui/visualisation/object_list_widget.py b/loopstructural/gui/visualisation/object_list_widget.py index 8d19021..9910852 100644 --- a/loopstructural/gui/visualisation/object_list_widget.py +++ b/loopstructural/gui/visualisation/object_list_widget.py @@ -35,6 +35,7 @@ def __init__(self, parent=None, *, viewer=None, properties_widget=None): self.setLayout(self.mainLayout) self.viewer = viewer self.viewer.objectAdded.connect(self.update_object_list) + self.viewer.outOfDateChanged.connect(self._on_out_of_date_changed) self.treeWidget.installEventFilter(self) self.treeWidget.itemSelectionChanged.connect(self.on_object_selected) self.treeWidget.itemDoubleClicked.connect(self.onDoubleClick) @@ -94,6 +95,9 @@ def update_object_list(self, new_object): mesh = meshes[mesh_name] self.add_mesh_item(mesh_name, mesh) + def _on_out_of_date_changed(self): + self.update_object_list(None) + def add_mesh_item(self, mesh_name, mesh): """Add a top-level tree item for a mesh and populate children for point/cell data arrays. @@ -139,7 +143,13 @@ def _on_vis(state, name=mesh_name, m=mesh): itemLayout = QHBoxLayout(itemWidget) itemLayout.setContentsMargins(0, 0, 0, 0) itemLayout.addWidget(visibilityCheckbox) - itemLayout.addWidget(QLabel(mesh_name)) + nameLabel = QLabel(mesh_name) + if isinstance(mesh, dict) and mesh.get('out_of_date'): + # the label text is the object name used elsewhere, so show the + # state with the style and tooltip only + nameLabel.setStyleSheet("color: gray; font-style: italic;") + nameLabel.setToolTip("Out of date: the model changed after this object was added") + itemLayout.addWidget(nameLabel) itemWidget.setLayout(itemLayout) self.treeWidget.setItemWidget(top, 0, itemWidget) diff --git a/loopstructural/gui/visualisation/object_properties_widget.py b/loopstructural/gui/visualisation/object_properties_widget.py index 562446c..550fafb 100644 --- a/loopstructural/gui/visualisation/object_properties_widget.py +++ b/loopstructural/gui/visualisation/object_properties_widget.py @@ -184,13 +184,8 @@ def set_opacity(self, value: float): pass # store in metadata if self.current_object_name in getattr(self.viewer, 'meshes', {}): - self.viewer.meshes[self.current_object_name].set( - 'kwargs', - { - **self.viewer.meshes[self.current_object_name].get('kwargs', {}), - 'opacity': value, - }, - ) + entry = self.viewer.meshes[self.current_object_name] + entry['kwargs'] = {**(entry.get('kwargs') or {}), 'opacity': value} except Exception: pass @@ -514,6 +509,7 @@ def _on_scalar_changed(self, scalar_name: str): opacity = old_kwargs.get('opacity', None) show_scalar_bar = self.scalar_bar_checkbox.isChecked() + source = self.viewer.get_source_metadata(self.current_object_name) try: self.viewer.remove_object(self.current_object_name) except Exception: @@ -528,11 +524,12 @@ def _on_scalar_changed(self, scalar_name: str): clim=clim, opacity=opacity, show_scalar_bar=show_scalar_bar, + **source, ) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get('mesh') except Exception: try: - self.viewer.add_mesh_object(mesh, name=self.current_object_name) + self.viewer.add_mesh_object(mesh, name=self.current_object_name, **source) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get('mesh') except Exception: pass @@ -645,6 +642,7 @@ def _on_colormap_changed(self, cmap: str): opacity = old_kwargs.get('opacity', None) show_scalar_bar = self.scalar_bar_checkbox.isChecked() + source = self.viewer.get_source_metadata(self.current_object_name) try: self.viewer.remove_object(self.current_object_name) except Exception: @@ -659,11 +657,12 @@ def _on_colormap_changed(self, cmap: str): clim=clim, opacity=opacity, show_scalar_bar=show_scalar_bar, + **source, ) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get('mesh') except Exception: try: - self.viewer.add_mesh_object(mesh, name=self.current_object_name) + self.viewer.add_mesh_object(mesh, name=self.current_object_name, **source) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get( 'mesh' ) From f2c377959b1efbf71337bba88c58ea9d1a8eb4bc Mon Sep 17 00:00:00 2001 From: lachlangrose Date: Wed, 30 Sep 2026 12:51:42 +0930 Subject: [PATCH 3/3] feat: filter viewer objects by value range or by unit Add a Filter (Threshold) section to the object properties panel. For the unit ids of a block model, cross section or coloured topography, it shows a check box for each unit (with its name and colour); other arrays use a value range, with an option to invert it. The full mesh stays stored, so the filter can be changed or cleared, and it stays when the object is updated from the model. The objects coloured by stratigraphy now store the unit names and colours, and the viewer gets replace_mesh_object to add an object again with the same settings. --- .../gui/visualisation/feature_list_widget.py | 88 +++--- .../visualisation/loop_pyvistaqt_wrapper.py | 74 ++++- .../gui/visualisation/mesh_scalar_utils.py | 65 +++++ .../visualisation/object_properties_widget.py | 260 +++++++++++++++++- loopstructural/main/model_manager.py | 13 + tests/unit/test_threshold_mesh.py | 74 +++++ 6 files changed, 519 insertions(+), 55 deletions(-) create mode 100644 tests/unit/test_threshold_mesh.py diff --git a/loopstructural/gui/visualisation/feature_list_widget.py b/loopstructural/gui/visualisation/feature_list_widget.py index 303ee4c..5bbba5f 100644 --- a/loopstructural/gui/visualisation/feature_list_widget.py +++ b/loopstructural/gui/visualisation/feature_list_widget.py @@ -846,6 +846,7 @@ def _on_topography_colour_finished(self, result): ids, colours = result mesh = self.viewer.meshes['topography_surface']['mesh'] self._set_stratigraphy_arrays(mesh.point_data, ids, colours) + metadata = {'coloured': True, **self._stratigraphy_metadata(colours)} self.viewer.add_mesh_object( mesh, name='topography_surface', @@ -859,7 +860,7 @@ def _on_topography_colour_finished(self, result): # keeps it a true, direct colour-for-colour match. lighting=False, source_type='topography_surface', - metadata={'coloured': True}, + metadata=metadata, ) logger.info("Coloured topography surface by stratigraphic column.") @@ -1067,6 +1068,7 @@ def _add_cross_section_mesh(self, result, source_type): # of how the section plane/line happens to be oriented. lighting=False, source_type=source_type, + metadata=self._stratigraphy_metadata(colours), ) logger.info(f"Added cross section '{self._pending_cross_section_name}'.") @@ -1170,7 +1172,10 @@ def _on_block_model_finished(self, result): show_scalar_bar=False, show_edges=False, source_type='block_model', - metadata={'ncells': [int(n) for n in np.asarray(mesh.dimensions) - 1]}, + metadata={ + 'ncells': [int(n) for n in np.asarray(mesh.dimensions) - 1], + **self._stratigraphy_metadata(colours), + }, ) logger.info(f"Added block model '{self._pending_block_model_name}'.") @@ -1194,10 +1199,21 @@ def _set_stratigraphy_arrays(data, ids, colours): data['stratigraphy'] = np.asarray(ids) data['colour'] = stratigraphic_ids_to_rgb(ids, colours) - def _colour_by_stratigraphy(self, points, data): + def _stratigraphy_metadata(self, colours): + """Unit names and colours, indexed by unit id, for the unit check + boxes of the object properties panel.""" + return { + 'unit_names': list(self.model_manager.get_stratigraphic_unit_names()), + 'unit_colours': list(colours), + } + + def _colour_by_stratigraphy(self, points, data, metadata): + """Colour a mesh by the stratigraphic unit at `points`, and store + the unit names and colours in `metadata`.""" ids = self.model_manager.evaluate_stratigraphy_on_points(points) colours = self.model_manager.get_stratigraphic_column_colours() self._set_stratigraphy_arrays(data, ids, colours) + metadata.update(self._stratigraphy_metadata(colours)) # Viewer objects that `_rebuild_object` can build again from the model, # by `source_type` (fold constraints use the 'fold_constraint_' prefix). @@ -1329,9 +1345,9 @@ def target(progress_callback): progress_callback(f"Updating {spec['name']} ({i + 1} of {len(specs)})...") try: mesh, overrides = self._rebuild_object(spec) - results.append((spec['name'], mesh, overrides, None)) + results.append((spec['name'], mesh, overrides, spec['metadata'], None)) except Exception as e: - results.append((spec['name'], None, {}, str(e))) + results.append((spec['name'], None, {}, None, str(e))) return results self._update_objects_thread, self._update_objects_worker, self._update_objects_progress = ( @@ -1352,8 +1368,9 @@ def _rebuild_object(self, spec): Runs on a background thread, so it must not touch the viewer. Returns (mesh, overrides), where overrides are viewer settings that - come from the model (e.g. a unit colour). Raises if the object cannot - be built. + come from the model (e.g. a unit colour). `spec['metadata']` is + updated with new values from the model (e.g. the unit names). Raises + if the object cannot be built. """ source_type = spec['source_type'] feature_name = spec['source_feature'] @@ -1391,17 +1408,17 @@ def _rebuild_object(self, spec): overrides['color'] = surface.colour elif source_type in ('cross_section_plane', 'cross_section_line'): mesh = spec['mesh'] - self._colour_by_stratigraphy(mesh.points, mesh.point_data) + self._colour_by_stratigraphy(mesh.points, mesh.point_data, metadata) elif source_type == 'block_model': bb = model.bounding_box mesh = build_block_model_mesh(bb.origin, bb.maximum, metadata['ncells']) - self._colour_by_stratigraphy(mesh.cell_centers().points, mesh.cell_data) + self._colour_by_stratigraphy(mesh.cell_centers().points, mesh.cell_data, metadata) elif source_type == 'topography_surface': xx, yy, zz = self.model_manager.sample_dem_grid() mesh = pv.StructuredGrid(xx, yy, zz) mesh['Elevation'] = mesh.points[:, 2] if metadata.get('coloured'): - self._colour_by_stratigraphy(mesh.points, mesh.point_data) + self._colour_by_stratigraphy(mesh.points, mesh.point_data, metadata) else: raise ValueError(f"Cannot update objects of type '{source_type}'") @@ -1426,7 +1443,7 @@ def _finish_update_objects_task(self): def _on_update_objects_finished(self, results): self._finish_update_objects_task() failed = [] - for name, mesh, overrides, error in results: + for name, mesh, overrides, metadata, error in results: entry = self.viewer.meshes.get(name) if entry is None: # removed from the viewer while the update ran @@ -1434,8 +1451,13 @@ def _on_update_objects_finished(self, results): if error is not None: failed.append(f"{name}: {error}") continue - if not self._replace_viewer_object(name, entry, mesh, overrides): - failed.append(f"{name}: cannot add the new object to the viewer") + try: + self.viewer.replace_mesh_object( + name, mesh, overrides, out_of_date=False, metadata=metadata + ) + except Exception as e: + logger.exception(f"Cannot add updated object '{name}' to the viewer") + failed.append(f"{name}: {e}") try: self.viewer.render() except Exception: @@ -1449,46 +1471,6 @@ def _on_update_objects_finished(self, results): "These objects were not updated and are still out of date:\n\n" + "\n".join(failed), ) - def _replace_viewer_object(self, name, entry, mesh, overrides) -> bool: - """Put `mesh` in the viewer in place of the object `name`, with the - same source values, viewer settings and visibility.""" - source = self.viewer.get_source_metadata(name) - source['out_of_date'] = False - kwargs = { - key: value - for key, value in (entry.get('kwargs') or {}).items() - if key not in source and key != 'name' - } - kwargs.update(overrides) - # a colour picked in the object properties panel has priority - user_colour = entry.get('color') - if user_colour is not None: - kwargs['color'] = user_colour - actor = entry.get('actor') - visible = bool(getattr(actor, 'visibility', True)) - - # pyvista replaces the actor with the same name, so the old object - # stays in the viewer if the new one cannot be added - try: - self.viewer.add_mesh_object(mesh, name=name, **source, **kwargs) - except Exception: - # e.g. a scalar array selected in the properties panel that the - # new mesh does not have; add it with the default colouring - for key in ('scalars', 'cmap', 'clim', 'rgb'): - kwargs.pop(key, None) - try: - self.viewer.add_mesh_object(mesh, name=name, **source, **kwargs) - except Exception: - logger.exception(f"Cannot add updated object '{name}' to the viewer") - return False - - new_entry = self.viewer.meshes.get(name, {}) - if user_colour is not None: - new_entry['color'] = user_colour - if not visible and new_entry.get('actor') is not None: - new_entry['actor'].visibility = False - return True - def _on_update_objects_error(self, traceback_text): self._finish_update_objects_task() self._refresh_update_objects_button() diff --git a/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py b/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py index 97fab5d..f0c6dfa 100644 --- a/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py +++ b/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py @@ -3,6 +3,8 @@ from pyvistaqt import QtInteractor from qgis.PyQt.QtCore import pyqtSignal +from .mesh_scalar_utils import threshold_mesh + class LoopPyVistaQTPlotter(QtInteractor): objectAdded = pyqtSignal(QtInteractor) # Signal to request deletion @@ -46,6 +48,7 @@ def add_mesh_object( isovalue: Optional[float] = None, metadata: Optional[Dict[str, Any]] = None, out_of_date: bool = False, + threshold: Optional[Dict[str, Any]] = None, **kwargs, ) -> None: """Add a mesh to the plotter. @@ -82,6 +85,11 @@ def add_mesh_object( example the number of blocks of a block model). out_of_date : bool True if the mesh no longer matches the model. + threshold : Optional[dict] + Show only the part of the mesh whose values are in a range (see + `mesh_scalar_utils.threshold_mesh`). The full mesh is still + stored, so the filter can be changed or removed later. Raises + ValueError if no cells are in the range. Returns ------- @@ -122,8 +130,14 @@ def add_mesh_object( # merge any extra kwargs (allow caller to override default choices) add_kwargs.update(kwargs) + display_mesh = mesh + if threshold: + display_mesh = threshold_mesh(mesh, threshold) + if display_mesh.n_cells == 0: + raise ValueError("No cells are in the filter range") + # attempt to add to the underlying pyvista plotter - actor = self.add_mesh(mesh, name=name, **add_kwargs) + actor = self.add_mesh(display_mesh, name=name, **add_kwargs) # store the mesh, actor and kwargs for future re-adds # persist source metadata so callers can find meshes created from model features @@ -136,6 +150,9 @@ def add_mesh_object( 'isovalue': isovalue, 'metadata': dict(metadata or {}), 'out_of_date': out_of_date, + 'threshold': dict(threshold) if threshold else None, + # the mesh shown in the viewer (the filtered part of `mesh`) + 'display_mesh': display_mesh, } self.objectAdded.emit(self) @@ -154,8 +171,63 @@ def get_source_metadata(self, name: str) -> Dict[str, Any]: 'isovalue': entry.get('isovalue'), 'metadata': entry.get('metadata'), 'out_of_date': bool(entry.get('out_of_date', False)), + 'threshold': entry.get('threshold'), } + def replace_mesh_object(self, name: str, mesh=None, overrides=None, **source_updates) -> None: + """Add the object `name` again, and keep its source values, viewer + settings (colour map, opacity, colour picked by the user, ...) and + visibility. + + Parameters + ---------- + name : str + Name of an object in the viewer. + mesh : optional + A new mesh for the object (e.g. built again from the model). If + None, the current mesh is used. + overrides : Optional[dict] + Viewer settings that replace the stored ones (e.g. a new unit + colour). A colour picked by the user still has priority. + **source_updates + Source values to change, e.g. `threshold=...` or + `out_of_date=False`. + + pyvista replaces the actor that has the same name, so if the new + object cannot be added, the old object stays and the error is raised. + """ + entry = self.meshes[name] + if mesh is None: + mesh = entry['mesh'] + source = self.get_source_metadata(name) + source.update(source_updates) + kwargs = { + key: value + for key, value in (entry.get('kwargs') or {}).items() + if key not in source and key != 'name' + } + kwargs.update(overrides or {}) + user_colour = entry.get('color') + if user_colour is not None: + kwargs['color'] = user_colour + actor = entry.get('actor') + visible = bool(getattr(actor, 'visibility', True)) + + try: + self.add_mesh_object(mesh, name=name, **source, **kwargs) + except Exception: + # e.g. a scalar array selected in the properties panel that the + # new mesh does not have; add it with the default colouring + for key in ('scalars', 'cmap', 'clim', 'rgb'): + kwargs.pop(key, None) + self.add_mesh_object(mesh, name=name, **source, **kwargs) + + new_entry = self.meshes[name] + if user_colour is not None: + new_entry['color'] = user_colour + if not visible and new_entry.get('actor') is not None: + new_entry['actor'].visibility = False + def set_out_of_date(self, names, out_of_date: bool = True) -> None: """Mark the named objects as out of date (or up to date).""" changed = False diff --git a/loopstructural/gui/visualisation/mesh_scalar_utils.py b/loopstructural/gui/visualisation/mesh_scalar_utils.py index 2af621e..be84b74 100644 --- a/loopstructural/gui/visualisation/mesh_scalar_utils.py +++ b/loopstructural/gui/visualisation/mesh_scalar_utils.py @@ -176,3 +176,68 @@ def apply_colormap_lut(mapper, cmap, clim=None, nan_color=(0.6, 0.6, 0.6, 1.0)): pass except Exception: pass + + +def filter_array_names(mesh): + """Return the names of the arrays of `mesh` that a threshold filter can + use: the single-component point arrays, and the single-component cell + arrays as `cell:` (the same naming as `get_scalar_values`). + """ + names = [] + for prefix, data in (('', 'point_data'), ('cell:', 'cell_data')): + arrays = getattr(mesh, data, None) or {} + for key in sorted(arrays.keys()): + values = np.asarray(arrays[key]) + if values.ndim == 1 and np.issubdtype(values.dtype, np.number): + names.append(f"{prefix}{key}") + return names + + +def threshold_mesh(mesh, threshold): + """Return the part of `mesh` whose values pass a filter. + + `threshold` is a dict with `scalars` (the array name, `cell:` for + a cell array) and one of: + + - `min`, `max`: keep the values in this range (the limits are + included); with `invert`, keep the values outside the range + - `values`: keep only these values (e.g. the ids of some units) + + For a point array, a cell is kept only if all its points pass. Raises + KeyError if the array is not on the mesh. + """ + name = threshold['scalars'] + preference = 'point' + if name.startswith('cell:'): + name = name.split(':', 1)[1] + preference = 'cell' + data = mesh.cell_data if preference == 'cell' else mesh.point_data + if name not in data: + raise KeyError(f"The object has no {preference} array '{name}'") + + if 'values' in threshold: + # threshold a 0/1 mask on a shallow copy, so the result has the same + # type as a range filter and the mesh of the caller does not change + mask = np.isin(np.asarray(data[name]), list(threshold['values'])).astype(float) + masked = mesh.copy(deep=False) + mask_data = masked.cell_data if preference == 'cell' else masked.point_data + mask_data['_filter_mask'] = mask + result = masked.threshold( + value=(0.5, 1.5), + scalars='_filter_mask', + preference=preference, + all_scalars=preference == 'point', + ) + for arrays in (result.cell_data, result.point_data): + if '_filter_mask' in arrays: + del arrays['_filter_mask'] + return result + + low, high = sorted((float(threshold['min']), float(threshold['max']))) + return mesh.threshold( + value=(low, high), + scalars=name, + preference=preference, + invert=bool(threshold.get('invert', False)), + all_scalars=preference == 'point', + ) diff --git a/loopstructural/gui/visualisation/object_properties_widget.py b/loopstructural/gui/visualisation/object_properties_widget.py index 550fafb..597a136 100644 --- a/loopstructural/gui/visualisation/object_properties_widget.py +++ b/loopstructural/gui/visualisation/object_properties_widget.py @@ -1,15 +1,22 @@ import matplotlib.pyplot as plt +import numpy as np # Add plotting imports for scalar histogram from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas from qgis.PyQt.QtCore import Qt +from qgis.PyQt.QtGui import QColor, QIcon, QPixmap from qgis.PyQt.QtWidgets import ( QCheckBox, QColorDialog, QComboBox, + QDoubleSpinBox, + QGroupBox, QHBoxLayout, QLabel, QLineEdit, + QListWidget, + QListWidgetItem, + QMessageBox, QPushButton, QSizePolicy, QSlider, @@ -17,7 +24,12 @@ QWidget, ) -from .mesh_scalar_utils import apply_colormap_lut, get_scalar_values, render_histogram +from .mesh_scalar_utils import ( + apply_colormap_lut, + filter_array_names, + get_scalar_values, + render_histogram, +) class ObjectPropertiesWidget(QWidget): @@ -110,6 +122,64 @@ def __init__(self, parent=None, *, viewer=None): self.hist_canvas.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Fixed) layout.addWidget(self.hist_canvas) + # Filter: show only the part of the object whose values are in a + # range, or (for the unit ids of a block model, cross section or + # topography) only the units that are checked + self.filter_group = QGroupBox("Filter (Threshold)") + filter_layout = QVBoxLayout(self.filter_group) + filter_layout.addWidget(QLabel("Array:")) + self.filter_array_combo = QComboBox() + self.filter_array_combo.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Fixed) + self.filter_array_combo.currentTextChanged.connect(self._on_filter_array_changed) + filter_layout.addWidget(self.filter_array_combo) + self.filter_range_widget = QWidget() + filter_range_widget_layout = QVBoxLayout(self.filter_range_widget) + filter_range_widget_layout.setContentsMargins(0, 0, 0, 0) + filter_range_layout = QHBoxLayout() + filter_range_layout.setSpacing(6) + filter_range_layout.addWidget(QLabel("Range:")) + self.filter_min = QDoubleSpinBox() + self.filter_max = QDoubleSpinBox() + for box in (self.filter_min, self.filter_max): + box.setRange(-1e12, 1e12) + box.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Fixed) + filter_range_layout.addWidget(box) + filter_range_widget_layout.addLayout(filter_range_layout) + self.filter_invert_checkbox = QCheckBox("Invert (show values outside the range)") + filter_range_widget_layout.addWidget(self.filter_invert_checkbox) + filter_layout.addWidget(self.filter_range_widget) + + self.filter_units_widget = QWidget() + filter_units_layout = QVBoxLayout(self.filter_units_widget) + filter_units_layout.setContentsMargins(0, 0, 0, 0) + filter_units_layout.addWidget(QLabel("Units to show:")) + self.filter_units_list = QListWidget() + self.filter_units_list.setMaximumHeight(160) + self.filter_units_list.itemChanged.connect(self._on_unit_check_changed) + filter_units_layout.addWidget(self.filter_units_list) + filter_units_buttons_layout = QHBoxLayout() + check_all_button = QPushButton("Check All") + check_all_button.clicked.connect(lambda: self._set_all_units_checked(True)) + uncheck_all_button = QPushButton("Uncheck All") + uncheck_all_button.clicked.connect(lambda: self._set_all_units_checked(False)) + filter_units_buttons_layout.addWidget(check_all_button) + filter_units_buttons_layout.addWidget(uncheck_all_button) + filter_units_layout.addLayout(filter_units_buttons_layout) + self.filter_units_widget.setVisible(False) + filter_layout.addWidget(self.filter_units_widget) + filter_buttons_layout = QHBoxLayout() + self.filter_apply_button = QPushButton("Apply Filter") + self.filter_apply_button.clicked.connect(self.apply_filter) + self.filter_clear_button = QPushButton("Clear Filter") + self.filter_clear_button.clicked.connect(self.clear_filter) + filter_buttons_layout.addWidget(self.filter_apply_button) + filter_buttons_layout.addWidget(self.filter_clear_button) + filter_layout.addLayout(filter_buttons_layout) + self.filter_status_label = QLabel("") + filter_layout.addWidget(self.filter_status_label) + self.filter_group.setEnabled(False) + layout.addWidget(self.filter_group) + # Surface Color surface_color_layout = QHBoxLayout() surface_color_layout.setSpacing(6) @@ -452,6 +522,194 @@ def setCurrentObject(self, object_name: str): except Exception: self._update_histogram(None) + self._load_filter_state(mesh_entry) + + def _load_filter_state(self, mesh_entry): + """Show the filter of the current object, or the full value range + of an array if the object has no filter.""" + names = filter_array_names(self.current_mesh) if self.current_mesh is not None else [] + threshold = mesh_entry.get('threshold') if isinstance(mesh_entry, dict) else None + self.filter_array_combo.blockSignals(True) + self.filter_array_combo.clear() + self.filter_array_combo.addItems(names) + self.filter_array_combo.blockSignals(False) + self.filter_group.setEnabled(bool(names)) + if threshold and threshold.get('scalars') in names: + self.filter_array_combo.blockSignals(True) + self.filter_array_combo.setCurrentText(threshold['scalars']) + self.filter_array_combo.blockSignals(False) + self._set_filter_range_to_data(threshold['scalars']) + if 'min' in threshold: + self.filter_min.setValue(float(threshold['min'])) + self.filter_max.setValue(float(threshold['max'])) + self.filter_invert_checkbox.setChecked(bool(threshold.get('invert', False))) + self._update_filter_mode(threshold['scalars'], threshold) + elif names: + # the unit ids of a block model or cross section are the most + # likely array to filter on + if 'cell:stratigraphy' in names: + self.filter_array_combo.setCurrentText('cell:stratigraphy') + elif 'stratigraphy' in names: + self.filter_array_combo.setCurrentText('stratigraphy') + self._set_filter_range_to_data(self.filter_array_combo.currentText()) + self.filter_invert_checkbox.setChecked(False) + self._update_filter_mode(self.filter_array_combo.currentText(), None) + self._update_filter_status() + + def _on_filter_array_changed(self, array_name: str): + if array_name: + self._set_filter_range_to_data(array_name) + entry = self.viewer.meshes.get(self.current_object_name) if self.viewer else None + threshold = entry.get('threshold') if entry else None + self._update_filter_mode(array_name, threshold) + + @staticmethod + def _is_unit_array(array_name: str) -> bool: + return array_name.split(':', 1)[-1] == 'stratigraphy' + + def _update_filter_mode(self, array_name: str, threshold): + """Show unit check boxes for a unit id array, and the range controls + for any other array.""" + unit_mode = self._is_unit_array(array_name) + self.filter_range_widget.setVisible(not unit_mode) + self.filter_apply_button.setVisible(not unit_mode) + self.filter_units_widget.setVisible(unit_mode) + if unit_mode: + self._populate_unit_list(array_name, threshold) + + def _populate_unit_list(self, array_name: str, threshold): + """Fill the unit check boxes with the units that are in the array. + + The unit names and colours come from the object's metadata (stored + when the object was coloured by the stratigraphic column), indexed + by unit id. + """ + from matplotlib.colors import to_hex + + entry = self.viewer.meshes.get(self.current_object_name, {}) if self.viewer else {} + metadata = entry.get('metadata') or {} + names = metadata.get('unit_names') or [] + colours = metadata.get('unit_colours') or [] + values = self._get_scalar_values(array_name) + ids = [int(v) for v in np.unique(values)] if values is not None else [] + checked = None + if threshold and threshold.get('scalars') == array_name and 'values' in threshold: + checked = {int(v) for v in threshold['values']} + + self.filter_units_list.blockSignals(True) + self.filter_units_list.clear() + for unit_id in ids: + if 0 <= unit_id < len(names): + label = str(names[unit_id]) + elif unit_id < 0: + label = "Outside all units" + else: + label = f"Unit {unit_id}" + item = QListWidgetItem(label) + item.setData(Qt.UserRole, unit_id) + item.setFlags(item.flags() | Qt.ItemIsUserCheckable) + item.setCheckState( + Qt.Checked if checked is None or unit_id in checked else Qt.Unchecked + ) + if 0 <= unit_id < len(colours): + try: + swatch = QPixmap(12, 12) + swatch.fill(QColor(to_hex(colours[unit_id]))) + item.setIcon(QIcon(swatch)) + except Exception: + pass + self.filter_units_list.addItem(item) + self.filter_units_list.blockSignals(False) + + def _checked_unit_ids(self): + ids, total = [], self.filter_units_list.count() + for i in range(total): + item = self.filter_units_list.item(i) + if item.checkState() == Qt.Checked: + ids.append(int(item.data(Qt.UserRole))) + return ids, total + + def _on_unit_check_changed(self, _item=None): + ids, total = self._checked_unit_ids() + if not ids: + # an object with no cells cannot be shown; keep the last filter + self.filter_status_label.setText("Check at least one unit to show") + return + if len(ids) == total: + self._set_threshold(None) + else: + self._set_threshold({'scalars': self.filter_array_combo.currentText(), 'values': ids}) + + def _set_all_units_checked(self, checked: bool): + self.filter_units_list.blockSignals(True) + for i in range(self.filter_units_list.count()): + self.filter_units_list.item(i).setCheckState(Qt.Checked if checked else Qt.Unchecked) + self.filter_units_list.blockSignals(False) + self._on_unit_check_changed() + + def _set_filter_range_to_data(self, array_name: str): + """Set the filter range to the full range of values of the array.""" + values = self._get_scalar_values(array_name) + if values is None: + return + values = np.asarray(values) + finite = values[np.isfinite(values)] if values.dtype.kind == 'f' else values + if finite.size == 0: + return + low, high = float(finite.min()), float(finite.max()) + # show integer arrays (e.g. unit ids) without decimals + integer = values.dtype.kind in 'iub' + for box in (self.filter_min, self.filter_max): + box.setDecimals(0 if integer else 4) + box.setSingleStep(1.0 if integer else max((high - low) / 100.0, 1e-4)) + self.filter_min.setValue(low) + self.filter_max.setValue(high) + + def apply_filter(self): + array_name = self.filter_array_combo.currentText() + if not array_name: + return + self._set_threshold( + { + 'scalars': array_name, + 'min': self.filter_min.value(), + 'max': self.filter_max.value(), + 'invert': self.filter_invert_checkbox.isChecked(), + } + ) + + def clear_filter(self): + self._set_threshold(None) + if self.filter_units_list.count(): + self.filter_units_list.blockSignals(True) + for i in range(self.filter_units_list.count()): + self.filter_units_list.item(i).setCheckState(Qt.Checked) + self.filter_units_list.blockSignals(False) + + def _set_threshold(self, threshold): + name = self.current_object_name + if not name or self.viewer is None or name not in self.viewer.meshes: + return + try: + self.viewer.replace_mesh_object(name, threshold=threshold) + except Exception as e: + QMessageBox.warning(self, "Filter", f"Cannot apply the filter:\n{e}") + return + try: + self.viewer.render() + except Exception: + pass + self._update_filter_status() + + def _update_filter_status(self): + entry = self.viewer.meshes.get(self.current_object_name) if self.viewer else None + if not entry or not entry.get('threshold'): + self.filter_status_label.setText("No filter") + return + total = getattr(entry.get('mesh'), 'n_cells', 0) + shown = getattr(entry.get('display_mesh'), 'n_cells', 0) + self.filter_status_label.setText(f"Showing {shown} of {total} cells") + def _on_scalar_changed(self, scalar_name: str): # update histogram preview immediately try: diff --git a/loopstructural/main/model_manager.py b/loopstructural/main/model_manager.py index c3aa7be..fb573fc 100644 --- a/loopstructural/main/model_manager.py +++ b/loopstructural/main/model_manager.py @@ -2000,6 +2000,19 @@ def get_stratigraphic_column_colours(self) -> list: colours.extend(unit.colour for unit in group.units) return colours + def get_stratigraphic_unit_names(self) -> list: + """Return unit names ordered to line up with `evaluate_model`'s ids. + + `names[i]` is the name of whichever unit `evaluate_model` labels `i` + (see `get_stratigraphic_column_colours`). + """ + if self.model is None or self.model.stratigraphic_column is None: + return [] + names = [] + for group in reversed(self.model.stratigraphic_column.get_groups()): + names.extend(unit.name for unit in group.units) + return names + def get_units_without_colour(self) -> list: """Return the names of units whose colour is missing or invalid. diff --git a/tests/unit/test_threshold_mesh.py b/tests/unit/test_threshold_mesh.py new file mode 100644 index 0000000..6b463ed --- /dev/null +++ b/tests/unit/test_threshold_mesh.py @@ -0,0 +1,74 @@ +"""Pytest tests for the threshold filter helpers in +loopstructural/gui/visualisation/mesh_scalar_utils.py. +""" + +import numpy as np +import pytest +import pyvista as pv + +from loopstructural.gui.visualisation.mesh_scalar_utils import ( + filter_array_names, + threshold_mesh, +) + + +@pytest.fixture +def block_model(): + grid = pv.ImageData(dimensions=(5, 5, 5)) + grid.cell_data['stratigraphy'] = np.arange(grid.n_cells) % 4 + grid.cell_data['colour'] = np.zeros((grid.n_cells, 3), dtype=np.uint8) + grid.point_data['z'] = grid.points[:, 2] + return grid + + +def test_filter_array_names_lists_single_component_arrays(block_model): + assert filter_array_names(block_model) == ['z', 'cell:stratigraphy'] + + +def test_threshold_keeps_cells_in_range(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'min': 1, 'max': 2}) + assert set(np.unique(result.cell_data['stratigraphy'])) == {1, 2} + + +def test_threshold_invert_keeps_cells_outside_range(block_model): + result = threshold_mesh( + block_model, {'scalars': 'cell:stratigraphy', 'min': 1, 'max': 2, 'invert': True} + ) + assert set(np.unique(result.cell_data['stratigraphy'])) == {0, 3} + + +def test_threshold_accepts_limits_in_either_order(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'min': 2, 'max': 1}) + assert set(np.unique(result.cell_data['stratigraphy'])) == {1, 2} + + +def test_threshold_on_point_array(block_model): + result = threshold_mesh(block_model, {'scalars': 'z', 'min': 0, 'max': 2}) + assert result.n_cells > 0 + assert result.points[:, 2].max() <= 2 + + +def test_threshold_missing_array_raises(block_model): + with pytest.raises(KeyError): + threshold_mesh(block_model, {'scalars': 'cell:missing', 'min': 0, 'max': 1}) + + +def test_threshold_values_keeps_only_those_cells(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'values': [0, 3]}) + assert set(np.unique(result.cell_data['stratigraphy'])) == {0, 3} + assert '_filter_mask' not in result.cell_data + assert '_filter_mask' not in block_model.cell_data + + +def test_threshold_values_on_point_array_keeps_arrays(): + grid = pv.ImageData(dimensions=(5, 5, 5)) + grid.point_data['stratigraphy'] = (grid.points[:, 2] >= 2).astype(int) + result = threshold_mesh(grid, {'scalars': 'stratigraphy', 'values': [1]}) + assert result.n_cells > 0 + assert set(np.unique(result.point_data['stratigraphy'])) == {1} + assert '_filter_mask' not in grid.point_data + + +def test_threshold_no_values_gives_empty_mesh(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'values': []}) + assert result.n_cells == 0