diff --git a/docs/development/usability-plan.md b/docs/development/usability-plan.md
index b3a9270..d0f3278 100644
--- a/docs/development/usability-plan.md
+++ b/docs/development/usability-plan.md
@@ -302,27 +302,27 @@ Acceptance: a new user can go from an empty project to a solved model with the
### Phase 4: Model step and direct interpolation
-- [ ] Replace "Initialize Model", "Solve Model" and "Update Model Data" with one
+- [x] Replace "Initialize Model", "Solve Model" and "Update Model Data" with one
primary button. Its text and action come from the model state.
-- [ ] Show the problems from all steps before the build.
-- [ ] Before the build, calculate all out-of-date derived data again: basal
+- [x] Show the problems from all steps before the build.
+- [x] Before the build, calculate all out-of-date derived data again: basal
contacts first, then calculated thicknesses. Run this in a background
task with progress. If the calculation fails, stop the build and show
the error.
-- [ ] With "Calculate from geology polygons", give the extracted contacts
+- [x] With "Calculate from geology polygons", give the extracted contacts
directly to the model. Update the project contacts layer for display
only.
-- [ ] Add the start choice: "Build from a geological map" or "Interpolate
+- [x] Add the start choice: "Build from a geological map" or "Interpolate
surfaces from constraints". Save the choice with the state.
-- [ ] Constraint list for each feature: source layer, constraint type, field
+- [x] Constraint list for each feature: source layer, constraint type, field
mapping, weight, Z source (layer Z, DEM or constant).
-- [ ] Add the constraint types: value, interface, gradient/normal, tangent,
+- [x] Add the constraint types: value, interface, gradient/normal, tangent,
inequality, pairwise inequality.
-- [ ] Show generated constraints as read-only rows. Add "Detach" to make a
+- [x] Show generated constraints as read-only rows. Add "Detach" to make a
generated feature editable.
-- [ ] Build and preview one feature: an isoline on the map canvas, or a surface
+- [x] Build and preview one feature: an isoline on the map canvas, or a surface
in the 3D view.
-- [ ] Implement "Add Fault" in the model step.
+- [x] Implement "Add Fault" in the model step.
Files: `geological_model_tab/*.py`, `layer_selection_table.py`,
`feature_details_panel/*.py`, `main/model_manager.py`.
@@ -361,9 +361,10 @@ model, the model uses basal contacts and thicknesses that match the new order.
done), or only a guide? Recommendation: only a guide. Show problems, but do
not lock steps.
2. Is the Data Conversion dialog part of step 1, or a separate tool?
-3. In direct mode, which constraint types does each interpolator (FDI, PLI,
- surfe) support? The UI must hide types that the selected interpolator does
- not support.
+3. ~~In direct mode, which constraint types does each interpolator (FDI, PLI,
+ surfe) support?~~ Resolved: all interpolators accept the same constraint
+ types. LoopStructural converts them as needed. The UI shows the same types
+ for every interpolator and does not hide any.
4. Is there a minimum QGIS version for the new widgets?
5. Must the Processing algorithms also use the derived-data record, or only
the dock? Recommendation: only the dock. Processing runs are single runs
diff --git a/loopstructural/gui/modelling/geological_model_tab/add_fault_dialog.py b/loopstructural/gui/modelling/geological_model_tab/add_fault_dialog.py
index 7d35d46..64aac59 100644
--- a/loopstructural/gui/modelling/geological_model_tab/add_fault_dialog.py
+++ b/loopstructural/gui/modelling/geological_model_tab/add_fault_dialog.py
@@ -1,33 +1,144 @@
-import os
+from qgis.PyQt.QtWidgets import (
+ QDialog,
+ QDialogButtonBox,
+ QDoubleSpinBox,
+ QFormLayout,
+ QGroupBox,
+ QHBoxLayout,
+ QLabel,
+ QLineEdit,
+ QVBoxLayout,
+)
-from qgis.PyQt.QtWidgets import QDialog
-from qgis.PyQt.uic import loadUi
+from ....main import parametric_fault
+
+
+def _spin(low, high, decimals=2, step=1.0):
+ spin = QDoubleSpinBox()
+ spin.setRange(low, high)
+ spin.setDecimals(decimals)
+ spin.setSingleStep(step)
+ return spin
class AddFaultDialog(QDialog):
- def __init__(self, parent=None):
+ """Ask for a fault that is given by numbers: a centre, a strike, a dip and a size.
+
+ Faults that have a trace on the map come from the fault layer in step 3.
+ The start values come from the bounding box of the model.
+ """
+
+ def __init__(self, parent=None, *, model_manager=None):
super().__init__(parent)
- ui_path = os.path.join(os.path.dirname(__file__), 'add_fault_dialog.ui')
- loadUi(ui_path, self)
- self.setWindowTitle('Add Fault Feature')
- # You can access widgets by their objectName from the .ui file
- # Example: self.strike_input, self.dip_input, etc.
+ self.model_manager = model_manager
+ self.setWindowTitle('Add Fault')
+
+ layout = QVBoxLayout(self)
+ name_form = QFormLayout()
+ self.name_input = QLineEdit()
+ name_form.addRow("Name:", self.name_input)
+ layout.addLayout(name_form)
+
+ orientation = QGroupBox("Orientation")
+ form = QFormLayout(orientation)
+ self.strike_input = _spin(0, 360)
+ self.strike_input.setToolTip("Degrees clockwise from north. The fault dips to the right.")
+ self.dip_input = _spin(0.01, 90)
+ self.pitch_input = _spin(-180, 180)
+ self.pitch_input.setToolTip(
+ "The angle of the slip in the fault plane, from the strike direction towards "
+ "the down-dip direction. 0 is strike-slip. 90 is dip-slip."
+ )
+ self.displacement_input = _spin(-1e9, 1e9)
+ self.displacement_input.setToolTip("The size of the slip, in model units.")
+ form.addRow("Strike (°):", self.strike_input)
+ form.addRow("Dip (°):", self.dip_input)
+ form.addRow("Pitch (°):", self.pitch_input)
+ form.addRow("Displacement:", self.displacement_input)
+ layout.addWidget(orientation)
+
+ position = QGroupBox("Position")
+ form = QFormLayout(position)
+ centre_row = QHBoxLayout()
+ self.centre_inputs = [_spin(-1e9, 1e9) for _ in range(3)]
+ for label, spin in zip("XYZ", self.centre_inputs):
+ centre_row.addWidget(QLabel(label))
+ centre_row.addWidget(spin)
+ form.addRow("Centre:", centre_row)
+ layout.addWidget(position)
+
+ size = QGroupBox("Size")
+ form = QFormLayout(size)
+ self.major_input = _spin(0.01, 1e9)
+ self.major_input.setToolTip("How far the fault extends along its strike.")
+ self.intermediate_input = _spin(0.01, 1e9)
+ self.intermediate_input.setToolTip("How far the fault extends down its dip.")
+ self.minor_input = _spin(0.01, 1e9)
+ self.minor_input.setToolTip("How far the fault changes the field, across the fault.")
+ form.addRow("Length along strike:", self.major_input)
+ form.addRow("Extent down dip:", self.intermediate_input)
+ form.addRow("Influence distance:", self.minor_input)
+ layout.addWidget(size)
+
+ self.problem_label = QLabel()
+ self.problem_label.setWordWrap(True)
+ layout.addWidget(self.problem_label)
+ self.button_box = QDialogButtonBox(
+ QDialogButtonBox.StandardButton.Ok | QDialogButtonBox.StandardButton.Cancel
+ )
+ self.button_box.accepted.connect(self.accept)
+ self.button_box.rejected.connect(self.reject)
+ layout.addWidget(self.button_box)
+
+ self._set_start_values()
+ for spin in (
+ self.strike_input,
+ self.dip_input,
+ self.pitch_input,
+ self.displacement_input,
+ self.major_input,
+ self.intermediate_input,
+ self.minor_input,
+ *self.centre_inputs,
+ ):
+ spin.valueChanged.connect(self._validate)
+ self.name_input.textChanged.connect(self._validate)
+ self._validate()
+
+ def _set_start_values(self):
+ model = getattr(self.model_manager, 'model', None)
+ box = getattr(model, 'bounding_box', None)
+ if box is None:
+ spec = parametric_fault.default_spec([0, 0, 0], [1000, 1000, 1000])
+ else:
+ spec = parametric_fault.default_spec(box.origin, box.maximum)
+ self.strike_input.setValue(spec['strike'])
+ self.dip_input.setValue(spec['dip'])
+ self.pitch_input.setValue(spec['pitch'])
+ self.displacement_input.setValue(spec['displacement'])
+ for spin, value in zip(self.centre_inputs, spec['centre']):
+ spin.setValue(value)
+ self.major_input.setValue(spec['major_axis'])
+ self.intermediate_input.setValue(spec['intermediate_axis'])
+ self.minor_input.setValue(spec['minor_axis'])
+
+ def _taken_names(self):
+ return self.model_manager.used_names() if self.model_manager is not None else set()
+
+ def _validate(self, *args):
+ found = parametric_fault.problems(self.get_fault_data(), self._taken_names())
+ self.problem_label.setText("\n".join(found))
+ self.button_box.button(QDialogButtonBox.StandardButton.Ok).setEnabled(not found)
def get_fault_data(self):
return {
+ 'name': self.name_input.text().strip(),
'strike': self.strike_input.value(),
'dip': self.dip_input.value(),
- 'centre': (
- self.centre_x_input.value(),
- self.centre_y_input.value(),
- self.centre_z_input.value(),
- ),
- 'ellipsoid_extents': (
- self.extent_x_input.value(),
- self.extent_y_input.value(),
- self.extent_z_input.value(),
- ),
- 'displacement': self.displacement_input.value(),
'pitch': self.pitch_input.value(),
- 'name': self.name_input.text(),
+ 'displacement': self.displacement_input.value(),
+ 'centre': tuple(spin.value() for spin in self.centre_inputs),
+ 'major_axis': self.major_input.value(),
+ 'intermediate_axis': self.intermediate_input.value(),
+ 'minor_axis': self.minor_input.value(),
}
diff --git a/loopstructural/gui/modelling/geological_model_tab/add_fault_dialog.ui b/loopstructural/gui/modelling/geological_model_tab/add_fault_dialog.ui
deleted file mode 100644
index 1d93ee4..0000000
--- a/loopstructural/gui/modelling/geological_model_tab/add_fault_dialog.ui
+++ /dev/null
@@ -1,141 +0,0 @@
-
-
- AddFaultDialog
-
-
- Add Fault Feature
-
-
- -
-
-
-
-
-
- Name
-
-
-
- -
-
-
-
-
- -
-
-
- Orientation
-
-
-
-
-
-
- Strike (°)
-
-
-
- -
-
-
- -
-
-
- Dip (°)
-
-
-
- -
-
-
- -
-
-
- Pitch (°)
-
-
-
- -
-
-
- -
-
-
- Displacement
-
-
-
- -
-
-
-
-
-
- -
-
-
- Position
-
-
-
-
-
-
- Centre (X, Y, Z)
-
-
-
- -
-
-
-
-
-
- -
-
-
- -
-
-
-
-
-
-
-
- -
-
-
- Ellipsoid Size
-
-
-
-
-
-
- Extents (X, Y, Z)
-
-
-
- -
-
-
-
-
-
- -
-
-
- -
-
-
-
-
-
-
-
- -
-
-
- QDialogButtonBox::Ok|QDialogButtonBox::Cancel
-
-
-
-
-
-
-
-
diff --git a/loopstructural/gui/modelling/geological_model_tab/feature_details_panel/_base.py b/loopstructural/gui/modelling/geological_model_tab/feature_details_panel/_base.py
index c528973..3bce57d 100644
--- a/loopstructural/gui/modelling/geological_model_tab/feature_details_panel/_base.py
+++ b/loopstructural/gui/modelling/geological_model_tab/feature_details_panel/_base.py
@@ -1,3 +1,4 @@
+import numpy as np
from LoopStructural.modelling.features import StructuralFrame
from LoopStructural.utils import normal_vector_to_strike_and_dip
from qgis.gui import QgsCollapsibleGroupBox, QgsMapLayerComboBox
@@ -6,17 +7,22 @@
QComboBox,
QDoubleSpinBox,
QFormLayout,
+ QHBoxLayout,
QLabel,
QMessageBox,
QPushButton,
QScrollArea,
+ QSpinBox,
QVBoxLayout,
QWidget,
)
from qgis.utils import plugins
from LoopStructural import getLogger
+from loopstructural.main import preview
+from ....background_task import finish_background_task, start_background_task
+from ....messages import push_info, push_warning
from ..bounding_box_widget import BoundingBoxWidget
from ..layer_selection_table import LayerSelectionTable
@@ -149,6 +155,7 @@ def __init__(self, parent=None, *, feature=None, model_manager=None, data_manage
name_validator=lambda: (True, ''), # Always valid in this context
)
table_layout = QVBoxLayout()
+ table_layout.addWidget(self._build_detach_widget())
table_layout.addWidget(self.layer_table)
self.view_constraint_data_button = QPushButton("View Data Used by Interpolator")
self.view_constraint_data_button.setToolTip(
@@ -170,11 +177,218 @@ def __init__(self, parent=None, *, feature=None, model_manager=None, data_manage
group_box = QgsCollapsibleGroupBox('Interpolator Settings')
group_box.setLayout(form_layout)
self.layout.addWidget(group_box)
+ self.layout.addWidget(self._build_preview_widget())
self.layout.addWidget(table_group_box)
# this will call the addMidBlock and addExportBlock methods
self.addMidBlock()
self.addExportBlock()
+ def _build_preview_widget(self):
+ """Return the group with the button that shows the isolines of this feature on the map."""
+ group = QgsCollapsibleGroupBox('Preview')
+ row = QHBoxLayout(group)
+ self.preview_levels_spin = QSpinBox()
+ self.preview_levels_spin.setRange(1, 100)
+ self.preview_levels_spin.setValue(preview.DEFAULT_LEVEL_COUNT)
+ self.preview_levels_spin.setPrefix("Lines: ")
+ self.preview_levels_spin.setToolTip("The number of isolines.")
+ self.preview_button = QPushButton("Preview on map")
+ self.preview_button.setToolTip(
+ "Solve this feature only, and add its isolines on the ground (the DEM) to the "
+ "project as a temporary layer. The other features are not solved."
+ )
+ self.preview_button.clicked.connect(self.preview_on_map)
+ row.addWidget(self.preview_levels_spin)
+ row.addWidget(self.preview_button, 1)
+ self._preview_layer_id = None
+ return group
+
+ def preview_on_map(self):
+ """Solve this feature and show its isolines on the map canvas.
+
+ The grid is read on the GUI thread (the DEM can be a raster layer). The
+ solve and the lines are made on a background thread.
+ """
+ manager = self.model_manager
+ name = self.feature.name
+ model = getattr(manager, 'model', None)
+ if model is None or model.bounding_box is None:
+ QMessageBox.warning(self, "Preview", "Build the model first.")
+ return
+ bounding_box = model.bounding_box
+ x, y, points = preview.grid_points(
+ bounding_box.origin,
+ bounding_box.maximum,
+ preview.DEFAULT_RESOLUTION,
+ getattr(manager, 'dem_function', None),
+ )
+ count = self.preview_levels_spin.value()
+
+ def target(progress_callback):
+ progress_callback("Solving the feature...")
+ values = np.asarray(manager.evaluate_feature_on_points(name, points), dtype=float)
+ if values.ndim != 1:
+ raise ValueError("The feature has no scalar field to draw.")
+ values = values.reshape(len(y), len(x))
+ progress_callback("Making the lines...")
+ levels = preview.default_levels(values, count)
+ return preview.isolines(x, y, values, levels), levels
+
+ self.preview_button.setEnabled(False)
+ self._preview_task = start_background_task(
+ self,
+ target,
+ title="Preview",
+ initial_label="Solving the feature...",
+ on_progress=self._on_preview_progress,
+ on_finished=self._on_preview_finished,
+ on_error=self._on_preview_error,
+ )
+
+ def _on_preview_progress(self, message):
+ try:
+ self._preview_task[2].setLabelText(message)
+ except Exception:
+ pass
+
+ def _end_preview_task(self):
+ finish_background_task(*self._preview_task)
+ self.preview_button.setEnabled(True)
+
+ def _on_preview_error(self, traceback_text):
+ self._end_preview_task()
+ lines = traceback_text.strip().splitlines()
+ QMessageBox.critical(
+ self, "Preview failed", lines[-1] if lines else "The preview stopped with an error."
+ )
+
+ def _on_preview_finished(self, result):
+ self._end_preview_task()
+ lines, levels = result
+ if not lines:
+ push_warning(
+ "Preview",
+ f"'{self.feature.name}' has no isolines on the map. The field has no range "
+ "or no values in the model area.",
+ )
+ return
+ self._show_preview_lines(lines)
+ # the feature is solved now, so the ticks of the feature list change
+ self.model_manager.notify('model_updated')
+ push_info(
+ "Preview",
+ f"{len(lines)} lines for '{self.feature.name}', from {levels.min():g} to {levels.max():g}.",
+ )
+
+ def _show_preview_lines(self, lines):
+ """Add the isolines as a temporary layer. An earlier preview of this feature is replaced."""
+ from qgis.core import (
+ QgsFeature,
+ QgsField,
+ QgsGeometry,
+ QgsPointXY,
+ QgsProject,
+ QgsVectorLayer,
+ )
+
+ from loopstructural.gui.compatibility import QVariantCompat
+
+ project = QgsProject.instance()
+ previous = getattr(self, '_preview_layer_id', None)
+ if previous and project.mapLayer(previous) is not None:
+ project.removeMapLayer(previous)
+ crs = self.data_manager.get_model_crs() if self.data_manager is not None else None
+ crs_text = crs.authid() if crs is not None and crs.isValid() else project.crs().authid()
+ layer = QgsVectorLayer(
+ f"LineString?crs={crs_text}", f"Preview: {self.feature.name}", "memory"
+ )
+ provider = layer.dataProvider()
+ provider.addAttributes([QgsField("value", QVariantCompat.Double)])
+ layer.updateFields()
+ features = []
+ for level, coordinates in lines:
+ feature = QgsFeature(layer.fields())
+ feature.setGeometry(
+ QgsGeometry.fromPolylineXY([QgsPointXY(px, py) for px, py in coordinates])
+ )
+ feature.setAttribute("value", level)
+ features.append(feature)
+ provider.addFeatures(features)
+ layer.updateExtents()
+ project.addMapLayer(layer)
+ self._preview_layer_id = layer.id()
+
+ def _build_detach_widget(self):
+ """Return the row that tells where the data of a generated feature comes from.
+
+ A generated feature is one that the stratigraphic column makes. Its
+ contact and orientation rows are read-only. The user can add rows to
+ them, or detach the feature to keep a fixed copy of the data.
+ """
+ self.detach_widget = QWidget()
+ row = QHBoxLayout(self.detach_widget)
+ row.setContentsMargins(0, 0, 0, 0)
+ self.detach_label = QLabel()
+ self.detach_label.setWordWrap(True)
+ self.detach_button = QPushButton()
+ self.detach_button.clicked.connect(self._toggle_detached)
+ row.addWidget(self.detach_label, 1)
+ row.addWidget(self.detach_button)
+ self._refresh_detach_widget()
+ return self.detach_widget
+
+ def _refresh_detach_widget(self):
+ manager = self.model_manager
+ name = self.feature.name
+ if manager is None or not manager.is_generated(name):
+ self.detach_widget.hide()
+ return
+ self.detach_widget.show()
+ if manager.is_detached(name):
+ self.detach_label.setText(
+ "Detached. The rows from the column are a fixed copy. They do not change "
+ "when the column or the contacts change."
+ )
+ self.detach_button.setText("Follow the column again")
+ self.detach_button.setToolTip(
+ "Use the data of the stratigraphic column again. The model needs a rebuild."
+ )
+ else:
+ self.detach_label.setText(
+ "The read-only rows come from the stratigraphic column. Rows that you "
+ "add are used together with them."
+ )
+ self.detach_button.setText("Detach")
+ self.detach_button.setToolTip(
+ "Keep a fixed copy of the data of this feature, so that it does not "
+ "follow the column."
+ )
+
+ def _toggle_detached(self):
+ manager = self.model_manager
+ name = self.feature.name
+ if manager.is_detached(name):
+ reply = QMessageBox.question(
+ self,
+ "Follow the column again",
+ f"'{name}' will use the data of the stratigraphic column again. "
+ "The fixed copy is removed. Continue?",
+ )
+ if reply != QMessageBox.StandardButton.Yes:
+ return
+ manager.attach_feature(name)
+ push_info("Feature", f"'{name}' follows the stratigraphic column again. Rebuild the model.")
+ elif manager.detach_feature(name):
+ push_info("Feature", f"'{name}' is detached. Its data is a fixed copy.")
+ else:
+ QMessageBox.warning(
+ self, "Cannot detach", f"'{name}' has no data from the column to keep."
+ )
+ return
+ self.data_manager._sync_processed_feature_data()
+ self.layer_table.restore_table_state()
+ self._refresh_detach_widget()
+
def addMidBlock(self):
"""Base mid block is intentionally empty now — bounding-box controls
were moved into the export/evaluation section so they appear alongside
diff --git a/loopstructural/gui/modelling/geological_model_tab/geological_model_tab.py b/loopstructural/gui/modelling/geological_model_tab/geological_model_tab.py
index 2c2a78b..3550252 100644
--- a/loopstructural/gui/modelling/geological_model_tab/geological_model_tab.py
+++ b/loopstructural/gui/modelling/geological_model_tab/geological_model_tab.py
@@ -1,3 +1,5 @@
+import html
+
from LoopStructural.modelling.features import FeatureType
from qgis.PyQt.QtCore import QObject, Qt, QThread, pyqtSignal, pyqtSlot
from qgis.PyQt.QtGui import QColor, QIcon, QPainter, QPen, QPixmap
@@ -16,7 +18,12 @@
QWidget,
)
+from ....main.derived_refresh import DerivedRefresh, DerivedRefreshError, names_to_refresh
from ....main.model_manager import ModelSolveCancelled
+from ....main.workflow_mode import DEFAULT_WORKFLOW_MODE, WORKFLOW_MODE_CONSTRAINTS
+from ...messages import push_info
+from ..steps import build_plan
+from .add_fault_dialog import AddFaultDialog
from .add_foliation_dialog import AddFoliationDialog
from .add_unconformity_dialog import AddUnconformityDialog
from .feature_details_panel import (
@@ -67,14 +74,14 @@ def _build_status_icon(color: str, *, filled: bool, mark: str = None) -> QIcon:
'solved': "Model status: solved",
'stale': (
"Model status: faults, stratigraphic column or input data changed — "
- "re-run Initialize Model"
+ "rebuild the model"
),
}
-# Solve Model only rebuilds interpolators for features that already exist; it
-# can't apply a fault topology edit (that requires re-running Initialize
-# Model, see GeologicalModelManager._on_fault_topology_changed), so it stays
-# disabled outside these two states.
+# Solving only rebuilds interpolators for features that already exist; it
+# can't apply a fault topology edit (that requires a build, see
+# GeologicalModelManager._on_fault_topology_changed), so the primary button
+# builds in the other states.
_SOLVABLE_STATES = {'initialized', 'solved'}
@@ -153,17 +160,18 @@ def __init__(self, parent=None, *, model_manager=None, data_manager=None):
splitter.setStretchFactor(1, 0) # Feature details panel
splitter.setOrientation(Qt.Orientation.Horizontal) # Add horizontal slider
- # Initialize / Solve Model buttons + a status summary of where the
- # model currently is: empty -> initialized (unsolved) -> solved.
- self.initializeModelButton = QPushButton("Initialize Model")
- self.solveModelButton = QPushButton("Solve Model")
- self.solveModelButton.setEnabled(False) # nothing to solve until initialized
+ # One primary button. Its text and action come from the model state
+ # (see build_plan.choose_primary_action). The status label shows where
+ # the model is: empty -> initialized (unsolved) -> solved.
+ self.primaryButton = QPushButton("Build model")
+ self.primaryButton.clicked.connect(self.on_primary_clicked)
+ self._primary_action = build_plan.PrimaryAction(build_plan.ACTION_BUILD, "Build model")
+ self._task_running = False
self.modelStatusLabel = QLabel(_MODEL_STATE_LABELS['empty'])
buttonRow = QHBoxLayout()
buttonRow.setContentsMargins(0, 0, 0, 0)
- buttonRow.addWidget(self.initializeModelButton)
- buttonRow.addWidget(self.solveModelButton)
+ buttonRow.addWidget(self.primaryButton)
buttonRow.addStretch(1)
buttonRow.addWidget(self.modelStatusLabel)
buttonRowWidget = QWidget()
@@ -173,36 +181,17 @@ def __init__(self, parent=None, *, model_manager=None, data_manager=None):
buttonRowWidget.setSizePolicy(QSizePolicy.Policy.Preferred, QSizePolicy.Policy.Fixed)
mainLayout.insertWidget(0, buttonRowWidget, 0)
- # Shown when an input layer changed after the model data was read
- # from it. "Update Model Data" puts the new data into the existing
- # features, so the changes made to them are kept (Initialize Model
- # builds every feature again).
- self.layerChangedLabel = QLabel()
- self.layerChangedLabel.setWordWrap(True)
- self.updateModelDataButton = QPushButton("Update Model Data")
- self.updateModelDataButton.setToolTip(
- "Read the changed layers again and put the new data into the current "
- "features. Solve Model then uses the new data."
- )
- self.updateModelDataButton.clicked.connect(self.update_model_data)
- layerChangedRow = QHBoxLayout()
- layerChangedRow.setContentsMargins(0, 0, 0, 0)
- layerChangedRow.addWidget(self.layerChangedLabel, 1)
- layerChangedRow.addWidget(self.updateModelDataButton)
- self.layerChangedWidget = QWidget()
- self.layerChangedWidget.setLayout(layerChangedRow)
- self.layerChangedWidget.setSizePolicy(
- QSizePolicy.Policy.Preferred, QSizePolicy.Policy.Fixed
- )
- mainLayout.insertWidget(1, self.layerChangedWidget, 0)
- self.layerChangedWidget.hide()
+ # The problems of all steps. They show before the build, so the user
+ # sees why a build can fail or can give a poor model.
+ self.problemsLabel = QLabel()
+ self.problemsLabel.setWordWrap(True)
+ self.problemsLabel.setTextFormat(Qt.TextFormat.RichText)
+ self.problemsLabel.setSizePolicy(QSizePolicy.Policy.Preferred, QSizePolicy.Policy.Fixed)
+ mainLayout.insertWidget(1, self.problemsLabel, 0)
+ self.problemsLabel.hide()
+ self._problems_provider = None
if self.data_manager is not None:
- self.data_manager.add_layer_data_changed_callback(self._refresh_changed_layers)
-
- # Action buttons
-
- self.initializeModelButton.clicked.connect(self.initialize_model)
- self.solveModelButton.clicked.connect(self.solve_model)
+ self.data_manager.add_layer_data_changed_callback(self.refresh_primary_action)
# Connect feature selection to update details panel
self.featureList.itemClicked.connect(self.on_feature_selected)
@@ -218,24 +207,37 @@ def __init__(self, parent=None, *, model_manager=None, data_manager=None):
def show_add_feature_menu(self, *args):
menu = QMenu(self)
- add_fault = menu.addAction("Add Fault (not yet implemented)")
- # Unlike Add Foliation/Add Unconformity, there's no model_manager entry
- # point yet for a parametric (strike/dip/centre) fault -- faults are
- # currently only created from trace data via update_fault_points. Keep
- # the menu entry visible (so it's discoverable) but disabled, rather
- # than silently accepting input and doing nothing with it.
- add_fault.setEnabled(False)
- add_fault.setToolTip("Adding a fault from parameters isn't implemented yet.")
+ add_fault = menu.addAction("Add Fault")
+ add_fault.setToolTip(
+ "Add a fault from a centre, a strike, a dip and a size. "
+ "Faults from a trace layer are in step 3."
+ )
add_foliaton = menu.addAction("Add Foliation")
add_unconformity = menu.addAction("Add Unconformity")
buttonPosition = self.sender().mapToGlobal(self.sender().rect().bottomLeft())
action = menu.exec(buttonPosition)
- if action == add_foliaton:
+ if action == add_fault:
+ self.open_add_fault_dialog()
+ elif action == add_foliaton:
self.open_add_foliation_dialog()
elif action == add_unconformity:
self.open_add_unconformity_dialog()
+ def open_add_fault_dialog(self):
+ dialog = AddFaultDialog(self, model_manager=self.model_manager)
+ if dialog.exec() != dialog.Accepted:
+ return
+ try:
+ self.model_manager.add_parametric_fault(dialog.get_fault_data())
+ except ValueError as err:
+ QMessageBox.critical(self, "Cannot add the fault", str(err))
+ return
+ push_info(
+ "Fault", f"'{dialog.get_fault_data()['name']}' is added. Rebuild the model to use it."
+ )
+ self.refresh_primary_action()
+
def open_add_foliation_dialog(self):
dialog = AddFoliationDialog(
self, data_manager=self.data_manager, model_manager=self.model_manager
@@ -250,39 +252,85 @@ def open_add_unconformity_dialog(self):
if dialog.exec() == dialog.Accepted:
pass
- def initialize_model(self):
- # Run update_model in a background thread to avoid blocking the UI.
- if not self.model_manager:
- return
- if self.data_manager is not None:
- if not self.data_manager.is_bounding_box_set():
- QMessageBox.critical(
- self,
- "Bounding box required",
- "Please set the bounding box before initializing the model.",
- )
- return
+ def set_problems_provider(self, provider):
+ """Set the function that gives the problems of all steps.
- # Validate model CRS
- if not self.data_manager.is_model_crs_valid():
- crs = self.data_manager.get_model_crs()
- if crs is None or not crs.isValid():
- msg = "Model CRS is not set or invalid. Please select a valid projected CRS in the Model Definition tab."
- else:
- # Safely get CRS description
- try:
- crs_desc = crs.description() or crs.authid() or "Unknown"
- except Exception:
- crs_desc = crs.authid() if hasattr(crs, 'authid') else "Unknown"
- msg = f"Model CRS must be projected (in meters), not geographic.\nSelected CRS: {crs_desc}\n\nPlease select a valid projected CRS in the Model Definition tab."
+ ``provider()`` returns ``(step_key, message)`` pairs. The dock sets it,
+ because the tab does not know the other steps.
+ """
+ self._problems_provider = provider
+ self.refresh_primary_action()
+
+ def _blocked_reason(self):
+ """Return why no build can start now, or None."""
+ if self.model_manager is None or self.data_manager is None:
+ return None
+ if not self.data_manager.is_bounding_box_set():
+ return "Set the bounding box in step 1."
+ if not self.data_manager.is_model_crs_valid():
+ return "Select a projected model CRS in step 1."
+ return None
+
+ def _workflow_mode(self):
+ return getattr(self.data_manager, 'workflow_mode', DEFAULT_WORKFLOW_MODE)
+
+ def _derived_to_refresh(self):
+ """Return the derived results that the build calculates again.
+
+ With "Interpolate surfaces from constraints", there is no column, so
+ there is nothing to calculate.
+ """
+ if self.data_manager is None or self._workflow_mode() == WORKFLOW_MODE_CONSTRAINTS:
+ return []
+ return names_to_refresh(self.data_manager)
- QMessageBox.critical(
- self,
- "Invalid Model CRS",
- msg,
- )
- return
+ def refresh_primary_action(self, *args, **kwargs):
+ """Show the action of the primary button and the problems of the steps."""
+ if self.model_manager is None:
+ return
+ changed = self.data_manager.get_changed_layers() if self.data_manager else []
+ self._primary_action = build_plan.choose_primary_action(
+ self.model_manager.model_state,
+ derived_out_of_date=self._derived_to_refresh(),
+ layers_changed=bool(changed),
+ blocked_reason=self._blocked_reason(),
+ )
+ action = self._primary_action
+ self.primaryButton.setText(action.text)
+ self.primaryButton.setToolTip(action.tooltip)
+ if not self._task_running:
+ self.primaryButton.setEnabled(action.enabled)
+
+ problems = self._problems_provider() if self._problems_provider is not None else []
+ if problems:
+ lines = "".join(f"
{html.escape(message)}" for _key, message in problems)
+ self.problemsLabel.setText(
+ f"Check before the build:"
+ )
+ self.problemsLabel.show()
+ else:
+ self.problemsLabel.hide()
+ def on_primary_clicked(self):
+ """Do the action of the primary button."""
+ if self.model_manager is None or self._task_running:
+ return
+ self.refresh_primary_action()
+ action = self._primary_action
+ if action.action == build_plan.ACTION_BUILD:
+ self.build_model()
+ elif action.action == build_plan.ACTION_SOLVE:
+ self.solve_model()
+
+ def build_model(self):
+ """Calculate the out-of-date derived data, make the features, and solve them.
+
+ The derived data is calculated first, because a build must not use
+ out-of-date contacts or thicknesses. If a calculation fails, the build
+ stops.
+ """
+ if not self.model_manager or not self._validate_before_build():
+ return
if not self._confirm_bounding_box_contains_data():
return
@@ -290,55 +338,123 @@ def initialize_model(self):
# build from the current layer data, not the data read before
# the layers changed
self.data_manager.reload_changed_layers()
+ # the rows that the user added to generated features
+ self.data_manager.sync_extra_constraints()
- self._run_model_task(
- lambda progress_callback: self.model_manager.update_model(
+ refresh = None
+ names = self._derived_to_refresh()
+ if names:
+ try:
+ refresh = DerivedRefresh(self.data_manager, self.model_manager, names)
+ except DerivedRefreshError as err:
+ QMessageBox.critical(self, "Cannot build the model", str(err))
+ return
+
+ if refresh is not None:
+ self._run_model_task(
+ lambda progress_callback: refresh.run(progress_callback),
+ title="Updating derived data",
+ initial_label="Calculating the data that comes from the column...",
+ cancellable=False,
+ on_success=lambda: self._after_derived_refresh(refresh),
+ )
+ else:
+ self._run_build_task()
+
+ def _after_derived_refresh(self, refresh):
+ """Put the new derived data in place, then make the features."""
+ try:
+ skipped = refresh.finish()
+ except Exception as err:
+ QMessageBox.critical(self, "Cannot build the model", f"{type(err).__name__}: {err}")
+ return
+ if skipped:
+ push_info(
+ "Thickness",
+ "These units keep the thickness that you typed: " + ", ".join(skipped) + ".",
+ )
+ self._run_build_task()
+
+ def _run_build_task(self):
+ def target(progress_callback):
+ self.model_manager.update_model(
notify_observers=False, progress_callback=progress_callback
- ),
- title="Updating Model",
- initial_label="Updating geological model...",
+ )
+ self.model_manager.update_all_features(
+ progress_callback=progress_callback, notify_observers=False
+ )
+
+ self._run_model_task(
+ target,
+ title="Building Model",
+ initial_label="Building geological model...",
)
- def update_model_data(self):
+ def _validate_before_build(self):
+ if self.data_manager is None:
+ return True
+ if not self.data_manager.is_bounding_box_set():
+ QMessageBox.critical(
+ self,
+ "Bounding box required",
+ "Please set the bounding box before building the model.",
+ )
+ return False
+ if not self.data_manager.is_model_crs_valid():
+ crs = self.data_manager.get_model_crs()
+ if crs is None or not crs.isValid():
+ msg = (
+ "Model CRS is not set or invalid. "
+ "Please select a valid projected CRS in step 1."
+ )
+ else:
+ # Safely get CRS description
+ try:
+ crs_desc = crs.description() or crs.authid() or "Unknown"
+ except Exception:
+ crs_desc = crs.authid() if hasattr(crs, 'authid') else "Unknown"
+ msg = (
+ "Model CRS must be projected (in meters), not geographic.\n"
+ f"Selected CRS: {crs_desc}\n\nPlease select a valid projected CRS in step 1."
+ )
+ QMessageBox.critical(self, "Invalid Model CRS", msg)
+ return False
+ return True
+
+ def _update_model_data(self):
"""Put the data of the changed input layers into the current
- features, without Initialize Model."""
- if self.data_manager is None or self.model_manager is None:
- return
+ features, without a build. Returns False if the caller must not solve now."""
try:
result = self.data_manager.refresh_model_data()
except Exception as e:
QMessageBox.critical(self, "Update model data failed", str(e))
- return
- self._refresh_model_status()
- if result['needs_initialize']:
- names = "\n".join(f" - {name}" for name in result['needs_initialize'])
- QMessageBox.information(
- self,
- "Initialize Model needed",
- "The new data for these features cannot be put into the current "
- f"model:\n{names}\n\n"
- "Run Initialize Model to use it. Initialize Model builds all "
- "features again.",
- )
-
- def _refresh_changed_layers(self):
- names = self.data_manager.get_changed_layers() if self.data_manager else []
- if not names:
- self.layerChangedWidget.hide()
- return
- self.layerChangedLabel.setText(
- "Input layers changed after the model data was read: " + ", ".join(names)
+ return False
+ self.refresh_primary_action()
+ if not result['needs_initialize']:
+ return True
+ names = "\n".join(f" - {name}" for name in result['needs_initialize'])
+ reply = QMessageBox.question(
+ self,
+ "Rebuild needed",
+ "The new data for these features cannot be put into the current "
+ f"model:\n{names}\n\n"
+ "Rebuild the model to use it? A rebuild makes all features again.",
)
- self.layerChangedWidget.show()
+ if reply == QMessageBox.StandardButton.Yes:
+ self.build_model()
+ return False
def solve_model(self):
- # Build/interpolate every feature already added to the model. Only
- # meaningful once Initialize Model has created some features, and not
- # while a fault topology edit is pending re-Initialize.
+ # Solve every feature already added to the model. Only meaningful
+ # once a build has created some features, and not while a fault
+ # topology edit is pending a build.
if not self.model_manager or self.model_manager.model_state not in _SOLVABLE_STATES:
return
if not self._confirm_bounding_box_contains_data():
return
+ if self.data_manager is not None and self.data_manager.get_changed_layers():
+ if not self._update_model_data():
+ return
self._run_model_task(
lambda progress_callback: self.model_manager.update_all_features(
progress_callback=progress_callback, notify_observers=False
@@ -378,7 +494,9 @@ def _confirm_bounding_box_contains_data(self):
)
return reply == QMessageBox.StandardButton.Yes
- def _run_model_task(self, target, *, title, initial_label):
+ def _run_model_task(
+ self, target, *, title, initial_label, cancellable=True, on_success=None
+ ):
"""Run `target(progress_callback)` on a background QThread with a
non-modal progress dialog, so the rest of QGIS stays usable. Both
Initialize Model and Solve Model share this: they disable each other
@@ -398,11 +516,18 @@ def _run_model_task(self, target, *, title, initial_label):
progress.setWindowModality(Qt.WindowModality.NonModal)
progress.setWindowTitle(title)
progress.setMinimumDuration(0)
- progress.canceled.connect(self._on_task_cancel_requested)
+ if cancellable:
+ progress.canceled.connect(self._on_task_cancel_requested)
+ else:
+ # the calculation of derived data cannot stop part way
+ progress.setCancelButton(None)
progress.show()
- self.initializeModelButton.setEnabled(False)
- self.solveModelButton.setEnabled(False)
+ self._task_running = True
+ self._task_failed = False
+ self._task_on_success = on_success
+ self._task_cancellable = cancellable
+ self.primaryButton.setEnabled(False)
# Only one task runs at a time (buttons are disabled above for the
# duration), so it's safe to stash the per-run state needed by the
@@ -461,6 +586,7 @@ def _on_task_progress(self, message, current, total):
@pyqtSlot()
def _on_task_finished(self):
+ on_success = None if self._task_failed else self._task_on_success
try:
# notify observers now on the GUI thread
try:
@@ -473,6 +599,8 @@ def _on_task_finished(self):
self._debug.log_error("Error notifying observer", e)
finally:
self._finish_task()
+ if on_success is not None:
+ on_success()
@pyqtSlot(str)
def _on_task_cancelled(self, message):
@@ -480,10 +608,12 @@ def _on_task_cancelled(self, message):
# right after this, from the worker's `finally`) run its normal
# cleanup/refresh -- the feature list will reflect whatever was
# actually built before the cancellation took effect.
+ self._task_failed = True
print(f"{self._task_title} cancelled: {message}")
@pyqtSlot(str, str)
def _on_task_error(self, reason, tb):
+ self._task_failed = True
try:
box = QMessageBox(self)
box.setIcon(QMessageBox.Icon.Critical)
@@ -497,8 +627,8 @@ def _on_task_error(self, reason, tb):
self._finish_task()
def _finish_task(self):
- self.initializeModelButton.setEnabled(True)
- self.solveModelButton.setEnabled(self.model_manager.model_state in _SOLVABLE_STATES)
+ self._task_running = False
+ self.refresh_primary_action()
try:
# QProgressDialog.close() emits canceled() itself (same as
# clicking the Cancel button), so disconnect first -- otherwise
@@ -566,7 +696,7 @@ def _status_icon(self, built):
def _refresh_model_status(self):
state = self.model_manager.model_state if self.model_manager is not None else 'empty'
self.modelStatusLabel.setText(_MODEL_STATE_LABELS.get(state, "Model status: unknown"))
- self.solveModelButton.setEnabled(state in _SOLVABLE_STATES)
+ self.refresh_primary_action()
def on_feature_selected(self, item):
feature_name = item.text(0)
@@ -663,6 +793,7 @@ def delete_feature(self, item):
feature_name = item.text(0)
# so that Initialize Model does not build it again
self.model_manager.remove_manual_foliation(feature_name)
+ self.model_manager.remove_parametric_fault(feature_name)
# Attempt to remove from the underlying model in a few ways
try:
# Try model's __delitem__ if supported
diff --git a/loopstructural/gui/modelling/geological_model_tab/layer_selection_table.py b/loopstructural/gui/modelling/geological_model_tab/layer_selection_table.py
index 826aa8c..0a8c3a5 100644
--- a/loopstructural/gui/modelling/geological_model_tab/layer_selection_table.py
+++ b/loopstructural/gui/modelling/geological_model_tab/layer_selection_table.py
@@ -1,5 +1,6 @@
from qgis.core import QgsMapLayerProxyModel
from qgis.gui import QgsFieldComboBox, QgsMapLayerComboBox
+from qgis.PyQt.QtCore import Qt
from qgis.PyQt.QtWidgets import (
QCheckBox,
QComboBox,
@@ -16,6 +17,7 @@
)
from loopstructural.gui.compatibility import configure_layer_combo
+from loopstructural.main import constraints
class LayerSelectionTable(QWidget):
@@ -173,7 +175,13 @@ def add_item_row(self):
def _create_type_combo(self):
"""Create type selection combo box."""
combo = QComboBox()
- combo.addItems(["Value", "Form Line", "Orientation", "Inequality"])
+ for layer_type in constraints.CONSTRAINT_TYPES:
+ combo.addItem(layer_type)
+ combo.setItemData(
+ combo.count() - 1,
+ constraints.DESCRIPTIONS[layer_type],
+ Qt.ItemDataRole.ToolTipRole,
+ )
return combo
def _create_select_layer_button(self, row, type_combo):
@@ -444,6 +452,114 @@ def _setup_type_specific_fields(self, layout):
self._setup_inequality_fields(layout)
elif self.layer_type == "Form Line":
self._setup_form_line_fields(layout)
+ elif self.layer_type == constraints.INTERFACE:
+ self._setup_interface_fields(layout)
+ elif self.layer_type in (constraints.GRADIENT_NORMAL, constraints.TANGENT):
+ self._setup_vector_fields(layout)
+ elif self.layer_type == constraints.PAIRWISE_INEQUALITY:
+ self._setup_pairwise_fields(layout)
+ self._setup_common_fields(layout)
+
+ def _field_combo(self, existing_key, allow_empty=False):
+ """Return a field combo of the layer, set to the field of an existing row."""
+ combo = QgsFieldComboBox()
+ if allow_empty:
+ combo.setAllowEmptyFieldName(True)
+ combo.setLayer(self.layer_combo.currentLayer())
+ self.layer_combo.layerChanged.connect(combo.setLayer)
+ if self.existing_data.get(existing_key):
+ combo.setField(self.existing_data[existing_key])
+ return combo
+
+ def _setup_interface_fields(self, layout):
+ """Setup fields for the interface type.
+
+ The group field is optional. Without it, each feature of the layer
+ (for example each line) is one surface.
+ """
+ form = QFormLayout()
+ self.group_field_combo = self._field_combo('group_field', allow_empty=True)
+ self.group_field_combo.setToolTip(
+ "Points with the same value in this field are on one surface. "
+ "If it is empty, each feature of the layer is one surface."
+ )
+ form.addRow("Group field (optional):", self.group_field_combo)
+ layout.addLayout(form)
+ self.field_combos = {'group_field': self.group_field_combo}
+
+ def _setup_vector_fields(self, layout):
+ """Setup fields for the gradient/normal and tangent types."""
+ form = QFormLayout()
+ self.field_combos = {}
+ if self.layer_type == constraints.GRADIENT_NORMAL:
+ self.vector_kind_combo = QComboBox()
+ self.vector_kind_combo.addItem("Gradient", constraints.KIND_GRADIENT)
+ self.vector_kind_combo.addItem("Normal", constraints.KIND_NORMAL)
+ index = self.vector_kind_combo.findData(
+ self.existing_data.get('vector_kind', constraints.KIND_GRADIENT)
+ )
+ self.vector_kind_combo.setCurrentIndex(max(index, 0))
+ self.vector_kind_combo.setToolTip(
+ "A gradient has a direction and a size. A normal has a direction only."
+ )
+ form.addRow("Vector:", self.vector_kind_combo)
+ self.field_combos['vector_kind'] = self.vector_kind_combo
+ for key, label in (
+ ('vector_x_field', "X component:"),
+ ('vector_y_field', "Y component:"),
+ ('vector_z_field', "Z component:"),
+ ):
+ combo = self._field_combo(key)
+ form.addRow(label, combo)
+ self.field_combos[key] = combo
+ layout.addLayout(form)
+
+ def _setup_pairwise_fields(self, layout):
+ """Setup fields for the pairwise inequality type."""
+ form = QFormLayout()
+ self.pair_field_combo = self._field_combo('pair_field')
+ self.pair_field_combo.setToolTip(
+ "A number for each group of points. The groups are ordered by this number."
+ )
+ form.addRow("Group number field:", self.pair_field_combo)
+ layout.addLayout(form)
+ self.field_combos = {'pair_field': self.pair_field_combo}
+
+ def _setup_common_fields(self, layout):
+ """Setup the weight and the source of Z. Every type has them."""
+ form = QFormLayout()
+ self.weight_spin = QDoubleSpinBox()
+ self.weight_spin.setDecimals(3)
+ self.weight_spin.setRange(0.001, 1000.0)
+ self.weight_spin.setSingleStep(0.1)
+ self.weight_spin.setValue(self.existing_data.get('weight', constraints.DEFAULT_WEIGHT))
+ self.weight_spin.setToolTip("How much the constraint counts in the interpolation.")
+ form.addRow("Weight:", self.weight_spin)
+
+ self.z_source_combo = QComboBox()
+ for source in constraints.Z_SOURCES:
+ self.z_source_combo.addItem(constraints.Z_LABELS[source], source)
+ self.z_source_combo.setCurrentIndex(
+ max(self.z_source_combo.findData(self.existing_data.get('z_source', constraints.Z_LAYER)), 0)
+ )
+ self.z_source_combo.setToolTip(
+ "Where the Z of the points comes from: the geometry of the layer, the DEM "
+ "of the project, or one number for all points."
+ )
+ form.addRow("Z source:", self.z_source_combo)
+
+ self.z_value_spin = QDoubleSpinBox()
+ self.z_value_spin.setDecimals(3)
+ self.z_value_spin.setRange(-1e9, 1e9)
+ self.z_value_spin.setValue(float(self.existing_data.get('z_value', 0.0)))
+ form.addRow("Z value:", self.z_value_spin)
+ layout.addLayout(form)
+
+ def update_z_value_state():
+ self.z_value_spin.setEnabled(self.z_source_combo.currentData() == constraints.Z_CONSTANT)
+
+ self.z_source_combo.currentIndexChanged.connect(update_z_value_state)
+ update_z_value_state()
def _setup_orientation_fields(self, layout):
"""Setup fields for orientation data type."""
@@ -673,7 +789,11 @@ def _on_accepted(self):
'layer': self.layer_combo.currentLayer(),
'layer_name': self._unique_layer_name(self.layer_combo.currentLayer()),
'type': self.layer_type,
+ 'weight': self.weight_spin.value(),
+ 'z_source': self.z_source_combo.currentData(),
}
+ if self.layer_data['z_source'] == constraints.Z_CONSTANT:
+ self.layer_data['z_value'] = self.z_value_spin.value()
# Add type-specific data
if self.layer_type == "Orientation":
@@ -713,6 +833,25 @@ def _on_accepted(self):
'dip_weight_spin'
].value()
+ elif self.layer_type == constraints.INTERFACE:
+ # the group field is optional
+ group_field = self.field_combos['group_field'].currentField()
+ if group_field:
+ self.layer_data['group_field'] = group_field
+
+ elif self.layer_type in (
+ constraints.GRADIENT_NORMAL,
+ constraints.TANGENT,
+ constraints.PAIRWISE_INEQUALITY,
+ ):
+ for key, combo in self.field_combos.items():
+ if key == 'vector_kind':
+ self.layer_data[key] = combo.currentData()
+ else:
+ self.layer_data[key] = combo.currentField()
+ if constraints.missing_fields(self.layer_data):
+ return
+
self.accept()
def get_layer_data(self):
diff --git a/loopstructural/gui/modelling/modelling_widget.py b/loopstructural/gui/modelling/modelling_widget.py
index 9b57d3d..0e12f65 100644
--- a/loopstructural/gui/modelling/modelling_widget.py
+++ b/loopstructural/gui/modelling/modelling_widget.py
@@ -8,7 +8,7 @@
QWidget,
)
-from loopstructural.gui.modelling.steps import checks
+from loopstructural.gui.modelling.steps import build_plan, checks
from loopstructural.gui.modelling.steps.header import DockHeader
from loopstructural.gui.modelling.steps.pages import (
DataStep,
@@ -90,8 +90,8 @@ def __init__(
self.footer_label.setWordWrap(True)
self.back_button = QPushButton("< Back", self)
self.next_button = QPushButton("Next >", self)
- self.back_button.clicked.connect(lambda _checked=False: self.go_to(self.current_index - 1))
- self.next_button.clicked.connect(lambda _checked=False: self.go_to(self.current_index + 1))
+ self.back_button.clicked.connect(lambda _checked=False: self.go_to(self._neighbour(-1)))
+ self.next_button.clicked.connect(lambda _checked=False: self.go_to(self._neighbour(1)))
footer = QHBoxLayout()
footer.addWidget(self.footer_label, 1)
footer.addWidget(self.back_button)
@@ -104,9 +104,14 @@ def __init__(
mainLayout.addLayout(footer)
self.step_bar.currentChanged.connect(self.go_to)
+ self._last_checks = {}
+ self.model_step.tab.set_problems_provider(self.problems)
+ if self.data_manager is not None:
+ self.data_manager.add_workflow_mode_callback(self._apply_workflow_mode)
self._timer = QTimer(self)
self._timer.setInterval(CHECK_INTERVAL_MS)
self._timer.timeout.connect(self.refresh_status)
+ self._apply_workflow_mode(getattr(self.data_manager, 'workflow_mode', None))
self.go_to(0)
@property
@@ -116,12 +121,50 @@ def current_index(self):
def go_to(self, index):
"""Show a step. The index is limited to the steps that exist."""
index = max(0, min(index, len(self.pages) - 1))
+ if index not in self._visible_indices():
+ # a hidden step: show the next visible step, or the one before it
+ after = self._neighbour(1, index)
+ index = after if after != index else self._neighbour(-1, index)
self.stack.setCurrentIndex(index)
self.step_bar.set_current_index(index)
- self.back_button.setEnabled(index > 0)
- self.next_button.setEnabled(index < len(self.pages) - 1)
+ self.back_button.setEnabled(self._neighbour(-1, index) != index)
+ self.next_button.setEnabled(self._neighbour(1, index) != index)
self.refresh_status()
+ def _visible_indices(self):
+ return [i for i, page in enumerate(self.pages) if self.step_bar.is_step_visible(page.key)]
+
+ def _neighbour(self, direction, index=None):
+ """Return the index of the next visible step, or ``index`` if there is none."""
+ index = self.current_index if index is None else index
+ candidates = [
+ i for i in self._visible_indices() if (i - index) * direction > 0
+ ]
+ if not candidates:
+ return index
+ return min(candidates, key=lambda i: abs(i - index))
+
+ def _apply_workflow_mode(self, mode):
+ """Show the steps that the start choice needs."""
+ mode = mode or getattr(self.data_manager, 'workflow_mode', None)
+ visible = build_plan.steps_for_mode(mode, [page.key for page in self.pages])
+ for page in self.pages:
+ self.step_bar.set_step_visible(page.key, page.key in visible)
+ self.data_step.set_workflow_mode(mode)
+ if self.stack.currentIndex() >= 0 and self.pages[self.stack.currentIndex()].key not in visible:
+ self.go_to(self.stack.currentIndex())
+ else:
+ self.refresh_status()
+
+ def problems(self):
+ """Return the problems of the steps that show, as ``(step_key, message)`` pairs."""
+ visible = self._visible_indices()
+ return build_plan.collect_problems(
+ (page.key, self._last_checks[page.key])
+ for i, page in enumerate(self.pages)
+ if i in visible and page.key in self._last_checks
+ )
+
def show_step(self, key):
"""Show the step with this key, for example `checks.STEP_VIEW`."""
self.go_to([page.key for page in self.pages].index(key))
@@ -133,7 +176,9 @@ def check_all(self):
def refresh_status(self):
"""Run the checks again and show the results in the bar and the footer."""
results = self.check_all()
+ self._last_checks = results
self.step_bar.set_checks(results)
+ self.model_step.tab.refresh_primary_action()
page = self.pages[self.current_index]
check = results[page.key]
if check.messages:
diff --git a/loopstructural/gui/modelling/steps/build_plan.py b/loopstructural/gui/modelling/steps/build_plan.py
new file mode 100644
index 0000000..34899a0
--- /dev/null
+++ b/loopstructural/gui/modelling/steps/build_plan.py
@@ -0,0 +1,111 @@
+"""The primary action of the model step.
+
+One button replaces "Initialize Model", "Solve Model" and "Update Model
+Data". The text and the action of the button come from the state of the model
+and from the derived data. This module does not import QGIS, so the unit tests
+can run it.
+"""
+
+from dataclasses import dataclass
+
+from loopstructural.main.workflow_mode import (
+ CONSTRAINT_MODE_HIDDEN_STEPS,
+ WORKFLOW_MODE_CONSTRAINTS,
+)
+
+# The actions of the primary button
+ACTION_BUILD = 'build' # make the features again, then solve them
+ACTION_SOLVE = 'solve' # solve the features that exist
+ACTION_NONE = 'none' # nothing can be done now
+
+@dataclass(frozen=True)
+class PrimaryAction:
+ """The text, the action and the enabled state of the primary button."""
+
+ action: str
+ text: str
+ enabled: bool = True
+ tooltip: str = ""
+
+
+def choose_primary_action(
+ model_state,
+ *,
+ derived_out_of_date=(),
+ layers_changed=False,
+ blocked_reason=None,
+) -> PrimaryAction:
+ """Return the action of the primary button.
+
+ Parameters
+ ----------
+ model_state : str
+ `GeologicalModelManager.model_state`: empty, initialized, solved or stale.
+ derived_out_of_date : iterable of str
+ Names of the derived results that the build calculates again.
+ layers_changed : bool
+ True if an input layer changed after the model data was read.
+ blocked_reason : str, optional
+ A problem that stops every build, for example no bounding box.
+ """
+ if blocked_reason:
+ return PrimaryAction(ACTION_NONE, "Build model", False, blocked_reason)
+ derived_out_of_date = list(derived_out_of_date)
+ if model_state == 'empty':
+ return PrimaryAction(
+ ACTION_BUILD, "Build model", tooltip="Make the features and solve them."
+ )
+ if model_state == 'stale' or derived_out_of_date:
+ if derived_out_of_date:
+ reason = "The data that comes from the column is out of date."
+ else:
+ reason = "The inputs changed after the model was built."
+ return PrimaryAction(
+ ACTION_BUILD,
+ "Rebuild model",
+ tooltip=f"{reason} Make the features again and solve them.",
+ )
+ if layers_changed:
+ return PrimaryAction(
+ ACTION_SOLVE,
+ "Update data and solve",
+ tooltip=(
+ "An input layer changed. Read the layers again, put the new data into "
+ "the features, and solve them."
+ ),
+ )
+ if model_state == 'initialized':
+ return PrimaryAction(ACTION_SOLVE, "Solve model", tooltip="Solve the features.")
+ return PrimaryAction(
+ ACTION_SOLVE, "Solve again", tooltip="Solve the features again with their current settings."
+ )
+
+
+def collect_problems(step_checks):
+ """Return the problems of all steps as ``(step_key, message)`` pairs.
+
+ Parameters
+ ----------
+ step_checks : iterable of (str, StepCheck)
+ The step keys and the results of their checks, in the order of the steps.
+ """
+ problems = []
+ seen = set()
+ for key, check in step_checks:
+ for message in check.problems:
+ # The same problem can show in more than one step
+ if message in seen:
+ continue
+ seen.add(message)
+ problems.append((key, message))
+ return problems
+
+
+def steps_for_mode(mode, all_steps):
+ """Return the step keys that the start choice shows.
+
+ With "Interpolate surfaces from constraints", the stratigraphy and fault
+ steps are hidden. The user can show them later.
+ """
+ hidden = CONSTRAINT_MODE_HIDDEN_STEPS if mode == WORKFLOW_MODE_CONSTRAINTS else ()
+ return [key for key in all_steps if key not in hidden]
diff --git a/loopstructural/gui/modelling/steps/checks.py b/loopstructural/gui/modelling/steps/checks.py
index 4606291..8159e0c 100644
--- a/loopstructural/gui/modelling/steps/checks.py
+++ b/loopstructural/gui/modelling/steps/checks.py
@@ -6,6 +6,7 @@
"""
from loopstructural.main import derived_data, layer_roles
+from loopstructural.main.workflow_mode import WORKFLOW_MODE_CONSTRAINTS
from .status import StepCheck
@@ -106,7 +107,9 @@ def check_model(data_manager, model_manager=None) -> StepCheck:
problems.append("Set the bounding box in step 1.")
if not data_manager.is_model_crs_valid():
problems.append("The model CRS must be a projected CRS (in metres).")
- problems.extend(_out_of_date_messages(data_manager))
+ # With constraints only, there is no column, so no result comes from it
+ if getattr(data_manager, 'workflow_mode', None) != WORKFLOW_MODE_CONSTRAINTS:
+ problems.extend(_out_of_date_messages(data_manager))
changed = data_manager.get_changed_layers()
if changed:
problems.append("Input layers changed after the model was built: " + ", ".join(changed))
diff --git a/loopstructural/gui/modelling/steps/pages.py b/loopstructural/gui/modelling/steps/pages.py
index 7c7ec72..847fc84 100644
--- a/loopstructural/gui/modelling/steps/pages.py
+++ b/loopstructural/gui/modelling/steps/pages.py
@@ -8,6 +8,7 @@
from qgis.gui import QgsCollapsibleGroupBox
from qgis.PyQt.QtCore import Qt, pyqtSignal
from qgis.PyQt.QtWidgets import (
+ QComboBox,
QHBoxLayout,
QLabel,
QMenu,
@@ -24,6 +25,8 @@
from loopstructural.gui.modelling.model_definition import ModelDefinitionTab
from loopstructural.gui.modelling.model_definition.fault_layers import FaultLayersWidget
+from loopstructural.main.workflow_mode import WORKFLOW_MODE_LABELS, WORKFLOW_MODES
+
from . import checks
@@ -76,10 +79,36 @@ def __init__(self, parent=None, **kwargs):
row.addWidget(QLabel("Select the area and the source layers.", self), 1)
row.addWidget(convert)
layout.addLayout(row)
+
+ # The start choice. The "constraints" choice hides steps 2 and 3.
+ self.mode_combo = QComboBox(self)
+ for mode in WORKFLOW_MODES:
+ self.mode_combo.addItem(WORKFLOW_MODE_LABELS[mode], mode)
+ self.mode_combo.setToolTip(
+ "Steps 2 and 3 make features and constraints from a geological map. "
+ "Interpolate from constraints hides them. Change the choice to show them again."
+ )
+ self.mode_combo.currentIndexChanged.connect(self._on_mode_changed)
+ mode_row = QHBoxLayout()
+ mode_row.addWidget(QLabel("Method:", self))
+ mode_row.addWidget(self.mode_combo, 1)
+ layout.addLayout(mode_row)
self.tab = ModelDefinitionTab(self, data_manager=self.data_manager)
layout.addWidget(self.tab, 1)
+ def set_workflow_mode(self, mode):
+ index = self.mode_combo.findData(mode)
+ if index >= 0 and index != self.mode_combo.currentIndex():
+ self.mode_combo.blockSignals(True)
+ self.mode_combo.setCurrentIndex(index)
+ self.mode_combo.blockSignals(False)
+
+ def _on_mode_changed(self, _index):
+ if self.data_manager is not None:
+ self.data_manager.set_workflow_mode(self.mode_combo.currentData())
+
+
class StratigraphyStep(StepPage):
"""Step 2: the stratigraphic column, and the results that come from the map."""
diff --git a/loopstructural/gui/modelling/steps/step_bar.py b/loopstructural/gui/modelling/steps/step_bar.py
index e8d96c0..811ddf1 100644
--- a/loopstructural/gui/modelling/steps/step_bar.py
+++ b/loopstructural/gui/modelling/steps/step_bar.py
@@ -79,6 +79,14 @@ def __init__(self, steps, parent=None):
def count(self):
return len(self._buttons)
+ def set_step_visible(self, key, visible):
+ """Show or hide the button of a step."""
+ self._buttons[self._keys.index(key)].setVisible(visible)
+
+ def is_step_visible(self, key):
+ # `isVisibleTo` does not depend on the visibility of the parent
+ return self._buttons[self._keys.index(key)].isVisibleTo(self)
+
def current_index(self):
return self._group.checkedId()
diff --git a/loopstructural/main/constraints.py b/loopstructural/main/constraints.py
new file mode 100644
index 0000000..951d317
--- /dev/null
+++ b/loopstructural/main/constraints.py
@@ -0,0 +1,186 @@
+"""The constraint types of a feature that the user builds from layers.
+
+Each row of the constraint list of a feature is a dictionary (see
+`LayerSelectionTable`). This module knows the types, the keys that each type
+needs, and how to turn the points of a layer into the data frame that
+LoopStructural reads. All interpolators (FDI, PLI and surfe) take the same
+types. LoopStructural converts them when an interpolator needs another form.
+
+This module does not import QGIS, so the unit tests can run it.
+"""
+
+import numpy as np
+import pandas as pd
+
+VALUE = 'Value'
+INTERFACE = 'Interface'
+ORIENTATION = 'Orientation'
+GRADIENT_NORMAL = 'Gradient/Normal'
+TANGENT = 'Tangent'
+FORM_LINE = 'Form Line'
+INEQUALITY = 'Inequality'
+PAIRWISE_INEQUALITY = 'Pairwise Inequality'
+
+# In the order of the type list of the user interface
+CONSTRAINT_TYPES = (
+ VALUE,
+ INTERFACE,
+ ORIENTATION,
+ GRADIENT_NORMAL,
+ TANGENT,
+ FORM_LINE,
+ INEQUALITY,
+ PAIRWISE_INEQUALITY,
+)
+
+DESCRIPTIONS = {
+ VALUE: "The scalar field has a known value at the points.",
+ INTERFACE: "The points are on one surface. The value of the surface is not known.",
+ ORIENTATION: "Strike and dip of the surface at the points.",
+ GRADIENT_NORMAL: "A vector at the points, from three fields (x, y, z). A gradient has "
+ "a direction and a size. A normal has a direction only.",
+ TANGENT: "A vector in the surface at the points, from three fields (x, y, z).",
+ FORM_LINE: "Lines on the surface. The line gives a constant value or the strike.",
+ INEQUALITY: "The value at the points is above a lower limit and below an upper limit.",
+ PAIRWISE_INEQUALITY: "Groups of points. The values of one group are ordered against the "
+ "values of the other groups.",
+}
+
+Z_LAYER = 'layer'
+Z_DEM = 'dem'
+Z_CONSTANT = 'constant'
+Z_SOURCES = (Z_LAYER, Z_DEM, Z_CONSTANT)
+Z_LABELS = {
+ Z_LAYER: "Z of the layer",
+ Z_DEM: "DEM",
+ Z_CONSTANT: "Constant",
+}
+
+KIND_GRADIENT = 'gradient'
+KIND_NORMAL = 'normal'
+
+DEFAULT_WEIGHT = 1.0
+
+# The keys of a row that name a field of the layer, for each type. A row
+# without these keys is not complete.
+REQUIRED_FIELDS = {
+ VALUE: ('value_field',),
+ INTERFACE: (),
+ ORIENTATION: ('strike_field', 'dip_field'),
+ GRADIENT_NORMAL: ('vector_x_field', 'vector_y_field', 'vector_z_field'),
+ TANGENT: ('vector_x_field', 'vector_y_field', 'vector_z_field'),
+ FORM_LINE: (),
+ INEQUALITY: ('lower_field', 'upper_field'),
+ PAIRWISE_INEQUALITY: ('pair_field',),
+}
+
+
+def missing_fields(layer_data):
+ """Return the keys that a row needs and does not have."""
+ return [key for key in REQUIRED_FIELDS.get(layer_data.get('type'), ()) if not layer_data.get(key)]
+
+
+def solver_for(layer_type):
+ """Return the solver that a type needs, or None for the default.
+
+ Inequality constraints need the ADMM solver.
+ """
+ return 'admm' if layer_type in (INEQUALITY, PAIRWISE_INEQUALITY) else None
+
+
+def sample_layer(sampler, layer_data, dem_function, default_use_z=False):
+ """Return the points of a row, with the Z that the row asks for.
+
+ Parameters
+ ----------
+ sampler : callable
+ Gives a data frame of points: ``sampler(df, dem_function, use_z)``.
+ layer_data : dict
+ The row. ``z_source`` is ``layer``, ``dem`` or ``constant``. For a
+ constant, ``z_value`` is the Z. A row with no ``z_source`` (made by an
+ older version) uses ``default_use_z``.
+ """
+ source = layer_data.get('z_source')
+ use_z = default_use_z if source is None else source == Z_LAYER
+ points = sampler(layer_data['df'], dem_function, use_z)
+ if source == Z_CONSTANT:
+ points = points.copy()
+ points['Z'] = float(layer_data.get('z_value', 0.0))
+ return points
+
+
+def add_weight(rows, points, layer_data):
+ """Add the weight of the row to ``rows``, a slice of ``points``.
+
+ Rows that already have a weight keep it. Without a weight in the row of the
+ constraint list, LoopStructural uses a weight of 1.
+ """
+ weight = layer_data.get('weight')
+ if weight is None or 'w' in rows:
+ return rows
+ rows = rows.copy()
+ rows['w'] = float(weight)
+ return rows
+
+
+def _vector_rows(points, layer_data, names):
+ fields = [layer_data[key] for key in ('vector_x_field', 'vector_y_field', 'vector_z_field')]
+ rows = points[['X', 'Y', 'Z']].copy()
+ for name, field in zip(names, fields):
+ rows[name] = pd.to_numeric(points[field], errors='coerce')
+ rows = rows.dropna(subset=list(names))
+ # A zero vector has no direction
+ zero = (rows[list(names)] == 0).all(axis=1)
+ return rows[~zero]
+
+
+def interface_rows(points, layer_data, feature_name, offset=0):
+ """Return the interface rows, and the offset for the next interface row.
+
+ With a group field, the points that have the same value are on one
+ surface. Without it, each feature of the layer (for example each line) is
+ one surface.
+ """
+ field = layer_data.get('group_field')
+ if field:
+ codes, _ = pd.factorize(points[field])
+ points = points[codes >= 0].copy()
+ interface = codes[codes >= 0].astype(float)
+ else:
+ interface = points['feature_id'].to_numpy(float)
+ count = len(np.unique(interface))
+ rows = points[['X', 'Y', 'Z']].copy()
+ # Different groups need different numbers, also between the rows of a feature
+ _, interface = np.unique(interface, return_inverse=True)
+ rows['interface'] = interface.astype(float) + offset
+ rows['feature_name'] = feature_name
+ return rows, offset + count
+
+
+def constraint_rows(points, layer_data, feature_name):
+ """Return the data frame rows of the types that this module builds.
+
+ Types: gradient/normal, tangent, pairwise inequality. Value, orientation,
+ form line and inequality are in `ModelManager._foliation_data`. Interface
+ rows come from `interface_rows`.
+
+ Returns
+ -------
+ pandas.DataFrame
+ The rows, with the weight of the row of the constraint list.
+ """
+ layer_type = layer_data['type']
+ if layer_type == GRADIENT_NORMAL:
+ kind = layer_data.get('vector_kind', KIND_GRADIENT)
+ names = ('nx', 'ny', 'nz') if kind == KIND_NORMAL else ('gx', 'gy', 'gz')
+ rows = _vector_rows(points, layer_data, names)
+ elif layer_type == TANGENT:
+ rows = _vector_rows(points, layer_data, ('tx', 'ty', 'tz'))
+ elif layer_type == PAIRWISE_INEQUALITY:
+ rows = points[['X', 'Y', 'Z']].copy()
+ rows['pair_id'] = pd.to_numeric(points[layer_data['pair_field']], errors='coerce')
+ rows = rows.dropna(subset=['pair_id'])
+ else:
+ raise ValueError(f"Unknown layer type: {layer_type}")
+ rows['feature_name'] = feature_name
+ return add_weight(rows, points, layer_data)
diff --git a/loopstructural/main/data_manager.py b/loopstructural/main/data_manager.py
index cea7c0d..fceda2d 100644
--- a/loopstructural/main/data_manager.py
+++ b/loopstructural/main/data_manager.py
@@ -33,6 +33,7 @@
from .layer_roles import LayerRoles
from .m2l_api import paint_stratigraphic_order
from .vectorLayerWrapper import qgsLayerToGeoDataFrame
+from .workflow_mode import WORKFLOW_MODE_MAP, WORKFLOW_MODES
def _lookup_colour_ramp(ramp_name):
@@ -169,6 +170,10 @@ def __init__(self, *, project=None, mapCanvas=None, logger=None):
# For each derived result, the inputs of the last run. See
# `derived_data` for how an out-of-date result is found.
self.derived = DerivedData()
+ # How the user builds the model: from a geological map, or by
+ # interpolating surfaces from constraints. See `gui.modelling.steps.build_plan`.
+ self.workflow_mode = WORKFLOW_MODE_MAP
+ self._workflow_mode_callbacks = []
self.derived.register(derived_data.BASAL_CONTACTS, self.basal_contacts_inputs)
self.derived.register(derived_data.THICKNESS, self.thickness_inputs)
self.derived.register(derived_data.STYLED_FIELDS, self.styled_fields_inputs)
@@ -301,6 +306,20 @@ def set_debug_manager(self, debug_manager):
"""Set the debug manager for the tools that the data manager can start."""
self.debug_manager = debug_manager
+ def set_workflow_mode(self, mode):
+ """Set the start choice of the user and tell the listeners if it changed."""
+ if mode not in WORKFLOW_MODES:
+ raise ValueError(f"Unknown workflow mode '{mode}'.")
+ if mode != self.workflow_mode:
+ self.workflow_mode = mode
+ for callback in list(self._workflow_mode_callbacks):
+ callback(mode)
+
+ def add_workflow_mode_callback(self, callback):
+ """Call ``callback(mode)`` each time the start choice changes."""
+ if callback not in self._workflow_mode_callbacks:
+ self._workflow_mode_callbacks.append(callback)
+
def _on_layer_role_changed(self, role, value):
"""A layer role changed: the derived results can be out of date."""
self.derived.refresh()
@@ -1326,6 +1345,35 @@ def get_changed_layers(self):
model data was last read from them."""
return sorted(self._watched_layers[i][0].name() for i in self._changed_layer_ids)
+ def sync_extra_constraints(self):
+ """Read the layers of the constraints that the user added to generated features.
+
+ A generated feature is one that the stratigraphic column makes. Its
+ table of data layers can have rows of the user. The model manager adds
+ them to the data of the column when it builds the feature. A feature
+ that the user added (a manual foliation) has its own build, so it is
+ not in this list.
+ """
+ if self._model_manager is None:
+ return
+ manual = self._model_manager.manual_foliations
+ generated = set(self._model_manager.generated_feature_names())
+ model_crs = self.get_model_crs()
+ extras = {}
+ for name, entries in self.feature_data.items():
+ if name in manual or name not in generated:
+ continue
+ rows = {}
+ for key, entry in entries.items():
+ if entry.get('processed') or entry.get('layer') is None:
+ continue
+ row = dict(entry)
+ row['df'] = qgsLayerToGeoDataFrame(entry['layer'], target_crs=model_crs)
+ rows[key] = row
+ if rows:
+ extras[name] = {'data': rows, 'use_z_coordinate': True}
+ self._model_manager.extra_constraints = extras
+
def reload_changed_layers(self):
"""Read the model data again from the input layers that changed.
@@ -1371,6 +1419,7 @@ def refresh_model_data(self):
if self._model_manager is None:
raise RuntimeError("Model manager is not set.")
self.reload_changed_layers()
+ self.sync_extra_constraints()
return self._model_manager.refresh_feature_data()
def get_layers_outside_bounding_box(self):
@@ -1505,12 +1554,28 @@ def get_stratigraphic_column(self):
def clear_stratigraphic_column(self):
self._stratigraphic_column.clear()
- def update_stratigraphy(self):
- """Update the foliation features in the model manager."""
+ def update_stratigraphy(self, basal_contacts=None, unit_name_field=None):
+ """Update the foliation features in the model manager.
+
+ Parameters
+ ----------
+ basal_contacts : geopandas.GeoDataFrame, optional
+ Contacts that the model uses in place of the basal contacts layer.
+ A build with "Calculate from geology polygons" gives the extracted
+ contacts here, so that the model does not read the project layer.
+ unit_name_field : str, optional
+ The unit name field of ``basal_contacts``.
+ """
self.logger(message="Updating stratigraphy...", log_level=4)
if self._model_manager is not None:
model_crs = self.get_model_crs()
- if self._basal_contacts is not None:
+ if basal_contacts is not None:
+ self._model_manager.update_contact_traces(
+ basal_contacts,
+ unit_name_field=unit_name_field,
+ use_z_coordinate=False,
+ )
+ elif self._basal_contacts is not None:
self._model_manager.update_contact_traces(
qgsLayerToGeoDataFrame(self._basal_contacts['layer'], target_crs=model_crs),
unit_name_field=self._basal_contacts['unitname_field'],
@@ -1599,6 +1664,7 @@ def _sync_processed_feature_data(self):
unit_data = self._model_manager.get_stratigraphy_entry(unit.name)
if not unit_data:
continue
+ suffix = 'detached' if self._model_manager.is_detached(group.name) else 'auto'
contact = unit_data.get('contact')
if (
contact is not None
@@ -1608,8 +1674,9 @@ def _sync_processed_feature_data(self):
self._add_processed_feature_row(
group.name,
self._basal_contacts.get('layer'),
- 'Contact (auto)',
+ f'Contact ({suffix})',
unit.name,
+ suffix,
)
orientations = unit_data.get('orientations')
if (
@@ -1620,8 +1687,9 @@ def _sync_processed_feature_data(self):
self._add_processed_feature_row(
group.name,
self._structural_orientations.get('layer'),
- 'Orientation (auto)',
+ f'Orientation ({suffix})',
unit.name,
+ suffix,
)
if self._fault_traces is not None:
@@ -1635,7 +1703,9 @@ def _sync_processed_feature_data(self):
fault_name,
)
- def _add_processed_feature_row(self, feature_name, layer, type_label, source_name):
+ def _add_processed_feature_row(
+ self, feature_name, layer, type_label, source_name, suffix='auto'
+ ):
"""Add a single read-only, workflow-derived row to `feature_data`.
Keyed on a string distinct from a plain layer name so a processed row
@@ -1644,7 +1714,7 @@ def _add_processed_feature_row(self, feature_name, layer, type_label, source_nam
"""
if layer is None:
return
- display_name = f"{source_name} ({layer.name()}, auto)"
+ display_name = f"{source_name} ({layer.name()}, {suffix})"
self.feature_data[feature_name][display_name] = {
'layer': layer,
'layer_name': display_name,
@@ -1701,6 +1771,7 @@ def reset(self):
self.layer_roles.clear()
self.thickness_sources.clear()
self.derived.clear()
+ self.set_workflow_mode(WORKFLOW_MODE_MAP)
self.set_dem_layer(None)
self.use_dem = True
@@ -1741,6 +1812,8 @@ def save_state(self, filepath):
self._model_manager.save_model(str(model_path))
state['model_file'] = model_path.name
state['manual_foliations'] = self._model_manager.manual_foliations_to_dict()
+ state['detached_features'] = self._model_manager.detached_to_dict()
+ state['parametric_faults'] = self._model_manager.parametric_faults_to_dict()
with open(path, 'w') as f:
json.dump(state, f, indent=2)
@@ -1773,6 +1846,8 @@ def load_state(self, filepath):
# the pickled model already has these features; this lets
# Initialize Model build them again
self._model_manager.manual_foliations_from_dict(state.get('manual_foliations', {}))
+ self._model_manager.detached_from_dict(state.get('detached_features', {}))
+ self._model_manager.parametric_faults_from_dict(state.get('parametric_faults', {}))
# the data was just read from the layers
self._changed_layer_ids.clear()
self.refresh_layer_watchers()
@@ -1856,6 +1931,7 @@ def to_dict(self):
'use_project_crs': self._use_project_crs,
'layer_roles': self.layer_roles.to_dict(),
'derived_data': self.derived.to_dict(),
+ 'workflow_mode': self.workflow_mode,
'thickness_sources': self.thickness_sources.to_dict(),
}
@@ -1876,6 +1952,9 @@ def _restore_derived_state(self, data):
self.layer_roles.contacts_source = layer_roles.CONTACTS_FROM_GEOLOGY
self.thickness_sources.from_dict(data.get('thickness_sources'))
self.derived.from_dict(data.get('derived_data'))
+ # A state file of an older version has no start choice
+ mode = data.get('workflow_mode')
+ self.set_workflow_mode(mode if mode in WORKFLOW_MODES else WORKFLOW_MODE_MAP)
def from_dict(self, data):
"""Load data from a dictionary."""
diff --git a/loopstructural/main/derived_refresh.py b/loopstructural/main/derived_refresh.py
new file mode 100644
index 0000000..66dde38
--- /dev/null
+++ b/loopstructural/main/derived_refresh.py
@@ -0,0 +1,316 @@
+"""Calculate the out-of-date derived data again before a model build.
+
+A model build must not use out-of-date basal contacts or thicknesses. This
+module finds the results to calculate, reads their inputs on the GUI thread,
+calculates them (on a background thread) and records the new inputs (on the
+GUI thread).
+
+The order is fixed: the basal contacts first, then the thicknesses, because
+the thickness calculation reads the contacts.
+
+With "Calculate from geology polygons", the contacts go directly to the model.
+The project layer of the contacts is updated afterwards, for display only.
+"""
+
+from typing import Callable, List, Optional
+
+from . import derived_data, layer_roles
+
+BASAL_CONTACT_UNIT_FIELD = 'basal_unit'
+
+
+class DerivedRefreshError(RuntimeError):
+ """A derived result could not be calculated. The build must stop."""
+
+
+def names_to_refresh(data_manager) -> List[str]:
+ """Return the derived results that a build calculates, in the order of calculation.
+
+ - Basal contacts: with "Calculate from geology polygons", when they are
+ out of date, or never calculated. With "Use a contacts layer", never,
+ because the layer is an input of the user.
+ - Thickness: only when it is out of date. A thickness that was never
+ calculated can be typed, so the build does not calculate it.
+ """
+ derived = data_manager.derived
+ names = []
+ if (
+ data_manager.layer_roles.contacts_source == layer_roles.CONTACTS_FROM_GEOLOGY
+ and data_manager.get_stratigraphic_unit_names()
+ and data_manager.get_layer_role(layer_roles.GEOLOGY) is not None
+ and data_manager.get_layer_role(layer_roles.GEOLOGY_UNIT_FIELD)
+ and derived.status(derived_data.BASAL_CONTACTS) != derived_data.CURRENT
+ ):
+ names.append(derived_data.BASAL_CONTACTS)
+ if derived.is_out_of_date(derived_data.THICKNESS):
+ names.append(derived_data.THICKNESS)
+ return names
+
+
+def _check_contacts_inputs(inputs):
+ if not inputs.get('geology') or not inputs.get('unit_field'):
+ raise DerivedRefreshError(
+ "The basal contacts cannot be calculated: select the geology layer and the "
+ "unit name field in step 1."
+ )
+
+
+class DerivedRefresh:
+ """One refresh of the out-of-date derived data.
+
+ Make the object on the GUI thread (it reads the layers and the settings).
+ Call `run` on the background thread. Call `finish` on the GUI thread.
+ """
+
+ def __init__(self, data_manager, model_manager, names=None):
+ self.data_manager = data_manager
+ self.model_manager = model_manager
+ self.names = list(names if names is not None else names_to_refresh(data_manager))
+ self._contacts = None
+ self._thickness = None
+ self._contacts_result = None
+ self._thickness_result = None
+ if derived_data.BASAL_CONTACTS in self.names:
+ self._contacts = self._read_contacts_job()
+ if derived_data.THICKNESS in self.names:
+ self._thickness = self._read_thickness_job()
+
+ def __bool__(self):
+ return bool(self.names)
+
+ # -- the reads on the GUI thread -------------------------------------
+
+ def _read_contacts_job(self):
+ dm = self.data_manager
+ settings = dm.get_widget_settings('basal_contacts_widget', {}) or {}
+ geology = dm.get_layer_role(layer_roles.GEOLOGY)
+ faults = (
+ dm.find_layer_by_name(settings['faults_layer'])
+ if settings.get('faults_layer')
+ else dm.get_layer_role(layer_roles.FAULT_TRACES)
+ )
+ unit_field = dm.get_layer_role(layer_roles.GEOLOGY_UNIT_FIELD)
+ inputs = dm.basal_contacts_inputs(
+ geology=geology,
+ unit_field=unit_field,
+ faults=faults,
+ ignore_units=settings.get('ignore_units', []),
+ override_units=settings.get('basal_override_units', []),
+ )
+ _check_contacts_inputs(inputs)
+ target_crs = dm.get_model_crs()
+ if target_crs is None or not target_crs.isValid():
+ target_crs = geology.crs()
+ return {
+ 'inputs': inputs,
+ 'params': dict(
+ geology=geology,
+ stratigraphic_order=dm.get_stratigraphic_unit_names(),
+ faults=faults,
+ ignore_units=list(settings.get('ignore_units', [])),
+ unit_name_field=unit_field,
+ target_crs=target_crs,
+ unit_colours=dm.get_stratigraphic_unit_colours(),
+ basal_override_units=list(settings.get('basal_override_units', [])),
+ debug_manager=getattr(dm, 'debug_manager', None),
+ ),
+ }
+
+ def _read_thickness_job(self):
+ dm = self.data_manager
+ settings = dm.get_widget_settings('thickness_calculator_widget', {}) or {}
+ calculator_type = settings.get('calculator_type')
+ if not calculator_type:
+ raise DerivedRefreshError(
+ "The thicknesses are out of date, but the thickness calculator has no settings. "
+ "Run Calculate thickness in step 2 one time."
+ )
+
+ def layer(key, role=None):
+ name = settings.get(key)
+ if name:
+ return dm.find_layer_by_name(name)
+ return dm.get_layer_role(role) if role else None
+
+ geology = layer('geology_layer', layer_roles.GEOLOGY)
+ structure = layer('structure_layer', layer_roles.STRUCTURE)
+ cross_sections = layer('cross_sections_layer')
+ unit_field = settings.get('unit_name_field') or dm.get_layer_role(
+ layer_roles.GEOLOGY_UNIT_FIELD
+ )
+ if geology is None or not unit_field:
+ raise DerivedRefreshError(
+ "The thicknesses cannot be calculated: select the geology layer and the "
+ "unit name field in step 1."
+ )
+ if calculator_type == 'StructuralPoint' and structure is None:
+ raise DerivedRefreshError(
+ "The thicknesses cannot be calculated: select the structure layer in step 1."
+ )
+ if calculator_type == 'AlongSection' and cross_sections is None:
+ raise DerivedRefreshError(
+ "The thicknesses cannot be calculated: select the cross-sections layer in "
+ "the thickness calculator."
+ )
+ contacts_layer = None
+ if dm.layer_roles.contacts_source == layer_roles.CONTACTS_FROM_LAYER:
+ contacts_layer = dm.get_layer_role(layer_roles.BASAL_CONTACTS)
+ elif settings.get('basal_contacts_layer'):
+ contacts_layer = dm.find_layer_by_name(settings['basal_contacts_layer'])
+ inputs = dm.thickness_inputs(
+ geology=geology,
+ unit_field=unit_field,
+ contacts_layer=contacts_layer,
+ calculator_type=calculator_type,
+ structure=structure,
+ cross_sections=cross_sections,
+ )
+ orientation_types = ("Dip Direction", "Strike")
+ index = settings.get('orientation_type_index', 0)
+ orientation_type = orientation_types[index] if 0 <= index < len(orientation_types) else None
+ params = dict(
+ calculator_type=calculator_type,
+ dtm=layer('dtm_layer', layer_roles.DEM),
+ geology=geology,
+ basal_contacts=contacts_layer,
+ sampling_frequency=settings.get('sampling_frequency', 200),
+ structure=structure,
+ cross_sections=cross_sections,
+ unit_name_field=unit_field,
+ dip_field=settings.get('dip_field') or 'DIP',
+ dipdir_field=settings.get('dipdir_field') or 'DIPDIR',
+ basal_contacts_unit_name=settings.get('basal_unit_field'),
+ max_line_length=(
+ settings.get('max_line_length')
+ if calculator_type in ("StructuralPoint", "InterpolatedStructure")
+ else None
+ ),
+ stratigraphic_order=dm.get_stratigraphic_unit_names(),
+ debug_manager=getattr(dm, 'debug_manager', None),
+ )
+ if orientation_type is not None:
+ params['orientation_type'] = orientation_type
+ return {'inputs': inputs, 'params': params}
+
+ # -- the calculation on the background thread ------------------------
+
+ def run(self, progress: Optional[Callable[[str, int, int], None]] = None):
+ """Calculate the results. Raise `DerivedRefreshError` if one fails."""
+ total = len(self.names)
+
+ def report(message, current):
+ if progress is not None:
+ progress(message, current, total)
+
+ for index, name in enumerate(self.names):
+ label = derived_data.DESCRIPTIONS[name].lower()
+ report(f"Calculating the {label} ({index + 1} of {total})...", index)
+
+ def updater(message, _label=label, _index=index):
+ report(f"Calculating the {_label}: {message}", _index)
+
+ try:
+ if name == derived_data.BASAL_CONTACTS:
+ self._run_contacts(updater)
+ elif name == derived_data.THICKNESS:
+ self._run_thickness(updater)
+ except DerivedRefreshError:
+ raise
+ except Exception as err:
+ raise DerivedRefreshError(
+ f"The {label} could not be calculated: {type(err).__name__}: {err}"
+ ) from err
+ report("Derived data is up to date.", total)
+
+ def _run_contacts(self, updater):
+ from .m2l_api import extract_basal_contacts
+
+ job = self._contacts
+ result = extract_basal_contacts(updater=updater, **job['params'])
+ contacts = result['basal_contacts']
+ if contacts is None or contacts.empty:
+ raise DerivedRefreshError(
+ "No basal contacts were found with the geology layer and the stratigraphic "
+ "column. Check the unit name field and the order of the column."
+ )
+ self._contacts_result = contacts
+ # The model reads the contacts directly. The project layer is only for display.
+ self.data_manager.update_stratigraphy(
+ basal_contacts=contacts, unit_name_field=BASAL_CONTACT_UNIT_FIELD
+ )
+
+ def _run_thickness(self, updater):
+ from .m2l_api import calculate_thickness
+
+ result = calculate_thickness(updater=updater, **self._thickness['params'])
+ if not isinstance(result, dict) or result.get('thicknesses') is None:
+ raise DerivedRefreshError("The thickness calculation gave no result.")
+ self._thickness_result = result
+
+ # -- the changes on the GUI thread ----------------------------------
+
+ def finish(self):
+ """Put the results into the data manager and record their inputs.
+
+ Returns
+ -------
+ list of str
+ The names of the units that keep a thickness that the user typed.
+ """
+ dm = self.data_manager
+ skipped = []
+ if self._contacts_result is not None:
+ self._update_contacts_layer(self._contacts_result)
+ dm.derived.record(derived_data.BASAL_CONTACTS, inputs=self._contacts['inputs'])
+ if self._thickness_result is not None:
+ values = thickness_values(self._thickness_result['thicknesses'])
+ _, skipped = dm.apply_calculated_thicknesses(values)
+ dm.derived.record(derived_data.THICKNESS, inputs=self._thickness['inputs'])
+ return skipped
+
+ def _update_contacts_layer(self, contacts):
+ """Show the contacts in the project. The model does not read this layer."""
+ from qgis.core import QgsProject
+
+ from .vectorLayerWrapper import addGeoDataFrameToproject
+
+ dm = self.data_manager
+ old = dm.get_layer_role(layer_roles.BASAL_CONTACTS)
+ layer = addGeoDataFrameToproject(contacts, "Basal contacts")
+ dm.apply_stratigraphic_colours_to_layer(layer, BASAL_CONTACT_UNIT_FIELD)
+ dm.set_basal_contacts(layer, unitname_field=BASAL_CONTACT_UNIT_FIELD)
+ # An earlier layer of the plugin is replaced. A layer of the user is not removed.
+ try:
+ if (
+ old is not None
+ and old.id() != layer.id()
+ and old.dataProvider().name() == 'memory'
+ and old.name() == "Basal contacts"
+ ):
+ QgsProject.instance().removeMapLayer(old.id())
+ except RuntimeError:
+ pass
+
+
+def thickness_values(thicknesses):
+ """Return unit name -> thickness from the table of the thickness calculator.
+
+ The median is used if it is there, then the mean. A unit with no result
+ (map2loop uses -1) is left out.
+ """
+ columns = getattr(thicknesses, 'columns', [])
+ if 'ThicknessMedian' in columns:
+ column = 'ThicknessMedian'
+ elif 'ThicknessMean' in columns:
+ column = 'ThicknessMean'
+ else:
+ return {}
+ values = {}
+ for _, row in thicknesses.iterrows():
+ name = row.get('name') or row.get('UNITNAME')
+ if not name:
+ continue
+ value = row.get(column)
+ if derived_data.thickness_is_set(value):
+ values[name] = float(value)
+ return values
diff --git a/loopstructural/main/model_manager.py b/loopstructural/main/model_manager.py
index 903b81c..3a79f3b 100644
--- a/loopstructural/main/model_manager.py
+++ b/loopstructural/main/model_manager.py
@@ -29,6 +29,7 @@
from LoopStructural import GeologicalModel
from loopstructural.toolbelt.preferences import PlgSettingsStructure
+from ..main import constraints, parametric_fault
from ..main.data_types import FaultEntry, StratigraphyEntry
from ..main.helpers import qgisAttributeIsNone
@@ -193,6 +194,16 @@ def __init__(self, debug_manager=None):
# features, so it builds these again from here; see
# `_build_manual_foliations`. Insertion order is the build order.
self.manual_foliations: Dict[str, dict] = {}
+ # Features that the stratigraphic column makes (one for each group of
+ # units). `detached` keeps a copy of the data of a feature that the
+ # user detached: the data no longer follows the column and the
+ # contacts. `extra_constraints` are the rows that the user added to a
+ # generated feature, in the form of `manual_foliations[name]`.
+ self.detached: Dict[str, pd.DataFrame] = {}
+ # Faults that the user gave by numbers (see `parametric_fault`), by name.
+ # They have no fault trace, so they are not in `faults`.
+ self.parametric_faults: Dict[str, dict] = {}
+ self.extra_constraints: Dict[str, dict] = {}
# Observers managed by Observable base class
self.dem_function = lambda x, y: 0
# internal flag to temporarily suppress notifications (used when
@@ -271,6 +282,9 @@ def reset(self):
self.faults = defaultdict(dict)
self.stratigraphy = defaultdict(dict)
self.manual_foliations = {}
+ self.detached = {}
+ self.extra_constraints = {}
+ self.parametric_faults = {}
self.dem_function = lambda x, y: 0
self._topology_dirty = False
self._data_dirty = False
@@ -1020,6 +1034,7 @@ def update_foliation_features(self):
groupname,
data=data,
force_constrained=True,
+ **self._extra_kwargs(groupname),
nelements=PlgSettingsStructure.interpolator_nelements,
npw=PlgSettingsStructure.interpolator_npw,
cpw=PlgSettingsStructure.interpolator_cpw,
@@ -1035,12 +1050,32 @@ def update_foliation_features(self):
# foliation features were rebuilt; let observers know
self._emit('foliation_features_updated')
- def _group_data(self, group, isovalues) -> Optional[pd.DataFrame]:
+ def _group_data(
+ self, group, isovalues, include_extra=True, use_detached=True
+ ) -> Optional[pd.DataFrame]:
"""Return the contact and orientation data of the units in `group`
as one data frame for its foliation, or None if there is no data.
- `isovalues` is `stratigraphic_column.get_isovalues()`.
+ `isovalues` is `stratigraphic_column.get_isovalues()`. A detached
+ group gives its kept data, not the data of the column. With
+ `include_extra`, the rows that the user added to the feature are in the
+ result. With `use_detached` False, the data of the column is used also
+ for a detached group.
"""
+ if use_detached and group.name in self.detached:
+ data = [self.detached[group.name].copy()]
+ else:
+ data = self._column_group_data(group, isovalues)
+ extra = self._extra_data(group.name) if include_extra else None
+ if extra is not None:
+ data.append(extra)
+ if len(data) == 0:
+ return None
+ return pd.concat(data, ignore_index=True)
+
+ def _column_group_data(self, group, isovalues) -> list:
+ """Return the data frames of the contacts and the orientations of the
+ units in `group`."""
data = []
for u in group.units:
val = isovalues[u.name]['value']
@@ -1059,9 +1094,108 @@ def _group_data(self, group, isovalues) -> Optional[pd.DataFrame]:
orientations['val'] = np.nan
orientations['feature_name'] = group.name
data.append(orientations)
- if len(data) == 0:
+ return data
+
+ def _extra_data(self, name) -> Optional[pd.DataFrame]:
+ """Return the rows that the user added to the generated feature `name`."""
+ spec = self.extra_constraints.get(name)
+ if not spec or not spec.get('data'):
return None
- return pd.concat(data, ignore_index=True)
+ try:
+ data, _kwargs = self._foliation_data(
+ name, spec['data'], AllSampler(), spec.get('use_z_coordinate', True)
+ )
+ except Exception as e:
+ if self._debug_manager is not None:
+ self._debug_manager.log(
+ f"Could not read the added constraints of '{name}': {e}", log_level=2
+ )
+ return None
+ return data
+
+ def _extra_kwargs(self, name) -> dict:
+ """Return the `create_and_add_foliation` arguments that the rows added to `name` need."""
+ spec = self.extra_constraints.get(name)
+ if not spec or not spec.get('data'):
+ return {}
+ kwargs = {}
+ for layer_data in spec['data'].values():
+ solver = constraints.solver_for(layer_data.get('type'))
+ if solver:
+ kwargs['solver'] = solver
+ return kwargs
+
+ # -- detach a generated feature ---------------------------------------
+
+ def generated_feature_names(self):
+ """Return the names of the features that the stratigraphic column makes."""
+ if self.stratigraphic_column is None:
+ return []
+ return [group.name for group in self.stratigraphic_column.get_groups()]
+
+ def is_generated(self, name) -> bool:
+ return name in self.generated_feature_names()
+
+ def is_detached(self, name) -> bool:
+ return name in self.detached
+
+ def detach_feature(self, name) -> bool:
+ """Keep the data of a generated feature, so that the user can edit it.
+
+ The data stays the same as it is now. It no longer changes when the
+ column, the contacts or the orientations change.
+
+ Returns
+ -------
+ bool
+ True if the feature is detached. False if it is not a generated
+ feature, or it has no data.
+ """
+ group = self._group_by_name(name)
+ if group is None:
+ return False
+ data = self._group_data(
+ group, self.stratigraphic_column.get_isovalues(), include_extra=False
+ )
+ if data is None:
+ return False
+ self.detached[name] = data.copy()
+ self._emit('model_updated')
+ return True
+
+ def attach_feature(self, name) -> bool:
+ """Make a detached feature follow the column again.
+
+ The model needs a build to use the data of the column.
+ """
+ if self.detached.pop(name, None) is None:
+ return False
+ self._data_dirty = True
+ self._emit('model_updated')
+ return True
+
+ def _group_by_name(self, name):
+ if self.stratigraphic_column is None:
+ return None
+ for group in self.stratigraphic_column.get_groups():
+ if group.name == name:
+ return group
+ return None
+
+ def detached_to_dict(self) -> dict:
+ """Return `detached` in a form that `json.dump` can write."""
+ return {
+ name: {
+ 'columns': {str(c): [_json_value(v) for v in data[c]] for c in data.columns}
+ }
+ for name, data in self.detached.items()
+ }
+
+ def detached_from_dict(self, detached: dict):
+ """Replace `detached` with the data from `detached_to_dict`."""
+ self.detached = {
+ name: pd.DataFrame(entry['columns']) for name, entry in (detached or {}).items()
+ }
def _strip_spurious_regions_from_domain_faults(self):
"""Work around a LoopStructural core gap that corrupts a domain
@@ -1252,8 +1386,68 @@ def update_fault_features(self):
cpw=PlgSettingsStructure.interpolator_cpw,
regularisation=PlgSettingsStructure.interpolator_regularisation,
)
+ self._build_parametric_faults()
self.apply_fault_abutting_relationships()
+ # -- faults that the user gives by numbers ---------------------------
+
+ def used_names(self) -> set:
+ """Return the names that a new feature or fault cannot use."""
+ names = {f.name for f in self.features()}
+ names |= set(self.faults) | set(self.parametric_faults) | set(self.manual_foliations)
+ return names
+
+ def add_parametric_fault(self, spec: dict):
+ """Add a fault that is given by a centre, a strike, a dip and a size.
+
+ The fault is built the next time the model is built, together with the
+ other faults, so that the features of the column use it. Until then
+ the model is stale.
+
+ Raises
+ ------
+ ValueError
+ If the spec has a problem, see `parametric_fault.problems`.
+ """
+ spec = parametric_fault.clean_spec(spec)
+ found = parametric_fault.problems(spec, self.used_names())
+ if found:
+ raise ValueError(" ".join(found))
+ self.parametric_faults[spec['name']] = spec
+ self._data_dirty = True
+ self._emit('model_updated')
+
+ def remove_parametric_fault(self, name: str):
+ """Stop building the parametric fault `name`. The model is stale until the next build."""
+ if self.parametric_faults.pop(name, None) is not None:
+ self._data_dirty = True
+
+ def _build_parametric_faults(self):
+ for name, spec in self.parametric_faults.items():
+ self._report_progress(f"Building fault '{name}'")
+ try:
+ self.model.create_and_add_fault(
+ name,
+ data=parametric_fault.frame_data(spec),
+ nelements=PlgSettingsStructure.interpolator_nelements,
+ npw=PlgSettingsStructure.interpolator_npw,
+ cpw=PlgSettingsStructure.interpolator_cpw,
+ regularisation=PlgSettingsStructure.interpolator_regularisation,
+ **parametric_fault.fault_arguments(spec),
+ )
+ except Exception as e:
+ raise ValueError(f"Could not build fault '{name}': {e}") from e
+
+ def parametric_faults_to_dict(self) -> dict:
+ """Return `parametric_faults` in a form that `json.dump` can write."""
+ return {name: parametric_fault.clean_spec(spec) for name, spec in self.parametric_faults.items()}
+
+ def parametric_faults_from_dict(self, faults: dict):
+ """Replace `parametric_faults` with the specs from `parametric_faults_to_dict`."""
+ self.parametric_faults = {
+ name: parametric_fault.spec_from_dict(spec) for name, spec in (faults or {}).items()
+ }
+
def _get_feature_by_name_or_none(self, name):
"""Non-raising counterpart to `GeologicalModel.get_feature_by_name`.
@@ -1460,7 +1654,12 @@ def update_model(
)
self._progress_callback = progress_callback
displacement_fault_count = len(set(self.faults) - set(self.fault_boundaries.values()))
- self._progress_total = displacement_fault_count + group_count + len(self.manual_foliations)
+ self._progress_total = (
+ displacement_fault_count
+ + len(self.parametric_faults)
+ + group_count
+ + len(self.manual_foliations)
+ )
self._progress_current = 0
dbg = getattr(self, '_debug_manager', None)
if dbg is not None:
@@ -1741,10 +1940,12 @@ def add_foliation(
Name for the new foliation feature.
data : dict
Mapping of layer identifiers to dicts describing each layer. Each
- layer dict must include a 'type' key (one of 'Orientation',
- 'Form Line', 'Value', 'Inequality') and the fields required by
- that type (e.g. 'strike_field', 'dip_field', 'value_field',
- 'form_line_constraint', ...).
+ layer dict must include a 'type' key (one of
+ `constraints.CONSTRAINT_TYPES`) and the fields required by that
+ type (e.g. 'strike_field', 'dip_field', 'value_field',
+ 'form_line_constraint', ...). Each layer dict can also have a
+ 'weight' and a 'z_source' ('layer', 'dem' or 'constant', with
+ 'z_value'); see `constraints`.
folded_feature_name : str or None
Optional name of a feature to which the foliation should be
associated/converted (currently unused in this helper).
@@ -1880,21 +2081,25 @@ def _foliation_data(self, name, data, sampler=AllSampler(), use_z_coordinate=Fal
# auto-synced rows from the map2loop workflow (type values
# like 'Contact (auto)') aren't meant for this manual path
continue
- if layer_data['type'] == 'Orientation':
- df = sampler(layer_data['df'], self.dem_function, use_z_coordinate)
+ layer_type = layer_data['type']
+ if layer_type not in constraints.CONSTRAINT_TYPES:
+ raise ValueError(f"Unknown layer type: {layer_type}")
+ df = constraints.sample_layer(sampler, layer_data, self.dem_function, use_z_coordinate)
+ if layer_type == 'Orientation':
df['strike'] = df[layer_data['strike_field']]
if layer_data.get('orientation_format') == 'Dip Direction':
df['strike'] = df['strike'] - 90
df['dip'] = df[layer_data['dip_field']]
df['feature_name'] = name
- dfs.append(df[['X', 'Y', 'Z', 'strike', 'dip', 'feature_name']])
- elif layer_data['type'] == 'Form Line':
- df = sampler(layer_data['df'], self.dem_function, use_z_coordinate)
+ rows = df[['X', 'Y', 'Z', 'strike', 'dip', 'feature_name']]
+ dfs.append(constraints.add_weight(rows, df, layer_data))
+ elif layer_type == 'Form Line':
df['feature_name'] = name
if layer_data.get('form_line_constraint') == 'strike':
df[['tx', 'ty', 'tz']] = _form_line_tangent_vectors(df)
df = df.dropna(subset=['tx', 'ty', 'tz'])
- dfs.append(df[['X', 'Y', 'Z', 'tx', 'ty', 'tz', 'feature_name']])
+ rows = df[['X', 'Y', 'Z', 'tx', 'ty', 'tz', 'feature_name']]
+ dfs.append(constraints.add_weight(rows, df, layer_data))
dip = layer_data.get('form_line_dip')
if dip is not None and not df.empty:
# The line's own direction is a precisely known hard
@@ -1919,24 +2124,32 @@ def _foliation_data(self, name, data, sampler=AllSampler(), use_z_coordinate=Fal
dip_df['w'] = layer_data.get('form_line_dip_weight', 0.1)
dfs.append(dip_df[['X', 'Y', 'Z', 'strike', 'dip', 'w', 'feature_name']])
else:
- df['interface'] = df['feature_id'].astype(float) + interface_offset
- interface_offset += df['feature_id'].nunique()
- dfs.append(df[['X', 'Y', 'Z', 'interface', 'feature_name']])
- elif layer_data['type'] == 'Value':
- df = sampler(layer_data['df'], self.dem_function, use_z_coordinate)
+ rows, interface_offset = constraints.interface_rows(
+ df, {}, name, interface_offset
+ )
+ dfs.append(constraints.add_weight(rows, df, layer_data))
+ elif layer_type == 'Interface':
+ rows, interface_offset = constraints.interface_rows(
+ df, layer_data, name, interface_offset
+ )
+ dfs.append(constraints.add_weight(rows, df, layer_data))
+ elif layer_type == 'Value':
df['val'] = df[layer_data['value_field']]
df['feature_name'] = name
- dfs.append(df[['X', 'Y', 'Z', 'val', 'feature_name']])
-
- elif layer_data['type'] == 'Inequality':
- df = sampler(layer_data['df'], self.dem_function, use_z_coordinate)
+ dfs.append(
+ constraints.add_weight(df[['X', 'Y', 'Z', 'val', 'feature_name']], df, layer_data)
+ )
+ elif layer_type == 'Inequality':
df['l'] = df[layer_data['lower_field']]
df['u'] = df[layer_data['upper_field']]
df['feature_name'] = name
- dfs.append(df[['X', 'Y', 'Z', 'l', 'u', 'feature_name']])
- kwargs['solver'] = 'admm'
+ rows = df[['X', 'Y', 'Z', 'l', 'u', 'feature_name']]
+ dfs.append(constraints.add_weight(rows, df, layer_data))
else:
- raise ValueError(f"Unknown layer type: {layer_data['type']}")
+ dfs.append(constraints.constraint_rows(df, layer_data, name))
+ solver = constraints.solver_for(layer_type)
+ if solver:
+ kwargs['solver'] = solver
return pd.concat(dfs, ignore_index=True), kwargs
def add_unconformity(
diff --git a/loopstructural/main/parametric_fault.py b/loopstructural/main/parametric_fault.py
new file mode 100644
index 0000000..c275b65
--- /dev/null
+++ b/loopstructural/main/parametric_fault.py
@@ -0,0 +1,168 @@
+"""A fault that the user gives by numbers: a centre, a strike, a dip and a size.
+
+Most faults come from the fault trace layer. A parametric fault has no
+trace. It is one ellipsoid in space. This module makes the vectors and the
+data that LoopStructural needs from the numbers of the user.
+
+This module does not import QGIS, so the unit tests can run it.
+"""
+
+import math
+from typing import Iterable, List, Optional, Tuple
+
+import numpy as np
+import pandas as pd
+
+FIELDS = (
+ 'name',
+ 'strike',
+ 'dip',
+ 'pitch',
+ 'displacement',
+ 'centre',
+ 'major_axis',
+ 'intermediate_axis',
+ 'minor_axis',
+)
+
+
+def plane_vectors(strike: float, dip: float, pitch: float) -> Tuple[np.ndarray, np.ndarray]:
+ """Return the normal vector and the slip vector of a fault plane.
+
+ Parameters
+ ----------
+ strike : float
+ Strike in degrees clockwise from north. The plane dips to the right of
+ the strike direction (right-hand rule).
+ dip : float
+ Dip in degrees, from 0 (flat) to 90 (vertical).
+ pitch : float
+ The angle of the slip in the plane, in degrees from the strike
+ direction towards the down-dip direction. 0 is a strike-slip fault and
+ 90 is a dip-slip fault.
+
+ Returns
+ -------
+ normal, slip : numpy.ndarray
+ Unit vectors (x east, y north, z up). The normal points up. The slip
+ vector is in the plane, so it is at right angles to the normal.
+ """
+ strike_r, dip_r, pitch_r = (math.radians(a) for a in (strike, dip, pitch))
+ dip_direction = strike_r + math.pi / 2
+ normal = np.array(
+ [
+ math.sin(dip_r) * math.sin(dip_direction),
+ math.sin(dip_r) * math.cos(dip_direction),
+ math.cos(dip_r),
+ ]
+ )
+ along_strike = np.array([math.sin(strike_r), math.cos(strike_r), 0.0])
+ down_dip = np.array(
+ [
+ math.cos(dip_r) * math.sin(dip_direction),
+ math.cos(dip_r) * math.cos(dip_direction),
+ -math.sin(dip_r),
+ ]
+ )
+ slip = math.cos(pitch_r) * along_strike + math.sin(pitch_r) * down_dip
+ return normal, slip / np.linalg.norm(slip)
+
+
+def default_spec(origin: Iterable[float], maximum: Iterable[float]) -> dict:
+ """Return the start values for a fault in the model area (the centre and the sizes)."""
+ origin = np.asarray(list(origin), dtype=float)
+ maximum = np.asarray(list(maximum), dtype=float)
+ size = maximum - origin
+ major = float(max(size[0], size[1]))
+ return {
+ 'name': '',
+ 'strike': 0.0,
+ 'dip': 90.0,
+ 'pitch': 90.0,
+ 'displacement': 100.0,
+ 'centre': tuple(float(v) for v in (origin + maximum) / 2),
+ 'major_axis': major,
+ 'intermediate_axis': major,
+ 'minor_axis': major / 2,
+ }
+
+
+def problems(spec: dict, taken_names: Iterable[str] = ()) -> List[str]:
+ """Return what is wrong with a fault, as text for the user. An empty list is a good fault."""
+ found = []
+ name = str(spec.get('name', '')).strip()
+ if not name:
+ found.append("Give the fault a name.")
+ elif name in set(taken_names):
+ found.append(f"The name '{name}' is already used by a feature or a fault.")
+ if not 0 < float(spec.get('dip', 0)) <= 90:
+ found.append("The dip must be above 0 and not more than 90 degrees.")
+ for key, label in (
+ ('major_axis', "The length along strike"),
+ ('intermediate_axis', "The extent down dip"),
+ ('minor_axis', "The influence distance"),
+ ):
+ if not float(spec.get(key) or 0) > 0:
+ found.append(f"{label} must be more than 0.")
+ centre = spec.get('centre')
+ if centre is None or len(centre) != 3 or not all(math.isfinite(float(v)) for v in centre):
+ found.append("The centre needs X, Y and Z.")
+ return found
+
+
+def frame_data(spec: dict) -> pd.DataFrame:
+ """Return the data frame with the one point that LoopStructural needs.
+
+ The point is the centre of the fault, with the value 0 of the fault
+ surface. The vectors and the sizes go to LoopStructural as separate
+ arguments, see `fault_arguments`.
+ """
+ x, y, z = (float(v) for v in spec['centre'])
+ return pd.DataFrame(
+ {
+ 'X': [x],
+ 'Y': [y],
+ 'Z': [z],
+ 'feature_name': [spec['name']],
+ 'val': [0.0],
+ 'coord': [0],
+ }
+ )
+
+
+def fault_arguments(spec: dict) -> dict:
+ """Return the arguments of ``GeologicalModel.create_and_add_fault`` for a fault.
+
+ ``name`` and ``data`` are not in the result: pass the name, and
+ ``frame_data(spec)`` as the data.
+ """
+ normal, slip = plane_vectors(spec['strike'], spec['dip'], spec.get('pitch', 90.0))
+ return {
+ 'displacement': float(spec['displacement']),
+ 'fault_normal_vector': normal,
+ 'fault_slip_vector': slip,
+ 'fault_center': np.array([float(v) for v in spec['centre']]),
+ 'major_axis': float(spec['major_axis']),
+ 'intermediate_axis': float(spec['intermediate_axis']),
+ 'minor_axis': float(spec['minor_axis']),
+ 'fault_dip': float(spec['dip']),
+ }
+
+
+def clean_spec(spec: dict) -> dict:
+ """Return the spec as plain numbers and text, safe for JSON."""
+ result = {key: spec[key] for key in FIELDS if key in spec}
+ result['name'] = str(result['name']).strip()
+ for key in ('strike', 'dip', 'pitch', 'displacement', 'major_axis', 'intermediate_axis',
+ 'minor_axis'):
+ if key in result:
+ result[key] = float(result[key])
+ result['centre'] = [float(v) for v in result['centre']]
+ return result
+
+
+def spec_from_dict(data: Optional[dict]) -> dict:
+ """Read a spec that was saved. The centre is a tuple again."""
+ spec = clean_spec(data)
+ spec['centre'] = tuple(spec['centre'])
+ return spec
diff --git a/loopstructural/main/preview.py b/loopstructural/main/preview.py
new file mode 100644
index 0000000..8451457
--- /dev/null
+++ b/loopstructural/main/preview.py
@@ -0,0 +1,106 @@
+"""The map preview of one feature: isolines of its scalar field.
+
+The scalar field is evaluated on a regular grid over the model area. Each
+point has the Z of the DEM, so the lines are the trace of the surfaces on the
+ground, as on a geological map.
+
+This module does not import QGIS, so the unit tests can run it.
+"""
+
+from typing import Callable, List, Optional, Tuple
+
+import numpy as np
+
+DEFAULT_RESOLUTION = 150
+DEFAULT_LEVEL_COUNT = 10
+
+
+def grid_points(
+ origin, maximum, resolution: int, dem_function: Optional[Callable] = None
+) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
+ """Return the grid of the preview.
+
+ Parameters
+ ----------
+ origin, maximum : sequence of float
+ The lower and the upper corner of the model area. Only X and Y are used.
+ resolution : int
+ The number of points along each side.
+ dem_function : callable, optional
+ ``dem_function(x, y)`` gives the Z of the ground. Without it, Z is 0.
+
+ Returns
+ -------
+ x, y : numpy.ndarray
+ The coordinates along each side, with shape ``(resolution,)``.
+ points : numpy.ndarray
+ The points, with shape ``(resolution * resolution, 3)``. The order is
+ row by row: the first ``resolution`` points have the first Y.
+ """
+ if resolution < 2:
+ raise ValueError("The resolution must be 2 or more.")
+ x = np.linspace(origin[0], maximum[0], resolution)
+ y = np.linspace(origin[1], maximum[1], resolution)
+ xx, yy = np.meshgrid(x, y)
+ if dem_function is None:
+ zz = np.zeros_like(xx)
+ else:
+ zz = np.vectorize(dem_function, otypes=[float])(xx, yy)
+ points = np.column_stack([xx.ravel(), yy.ravel(), zz.ravel()])
+ return x, y, points
+
+
+def default_levels(values, count: int = DEFAULT_LEVEL_COUNT) -> np.ndarray:
+ """Return ``count`` levels, spaced evenly between the lowest and the highest value.
+
+ The lowest and the highest value are not levels, because a line there is a
+ point or an edge. An empty array means that the field has no range.
+ """
+ values = np.asarray(values, dtype=float)
+ values = values[np.isfinite(values)]
+ if values.size == 0 or count < 1:
+ return np.array([])
+ low, high = values.min(), values.max()
+ if not high > low:
+ return np.array([])
+ return np.linspace(low, high, count + 2)[1:-1]
+
+
+def isolines(x, y, values, levels) -> List[Tuple[float, np.ndarray]]:
+ """Return the lines of the field ``values`` at each level.
+
+ Parameters
+ ----------
+ x, y : array_like
+ The coordinates along each side of the grid.
+ values : array_like
+ The field, with shape ``(len(y), len(x))``. A value that is not finite
+ is a gap: no line crosses it.
+ levels : iterable of float
+ The levels of the lines.
+
+ Returns
+ -------
+ list of tuple
+ ``(level, coordinates)`` for each line. ``coordinates`` has shape
+ ``(n, 2)`` with X and Y.
+ """
+ try:
+ import contourpy
+ except ImportError as err: # pragma: no cover - contourpy comes with matplotlib
+ raise RuntimeError("The preview needs the contourpy package (part of matplotlib).") from err
+ values = np.asarray(values, dtype=float)
+ if values.shape != (len(y), len(x)):
+ raise ValueError("values must have the shape (len(y), len(x)).")
+ generator = contourpy.contour_generator(
+ np.asarray(x, dtype=float),
+ np.asarray(y, dtype=float),
+ np.ma.masked_invalid(values),
+ )
+ lines = []
+ for level in levels:
+ for line in generator.lines(float(level)):
+ line = np.asarray(line, dtype=float)
+ if len(line) >= 2:
+ lines.append((float(level), line))
+ return lines
diff --git a/loopstructural/main/workflow_mode.py b/loopstructural/main/workflow_mode.py
new file mode 100644
index 0000000..db1aea1
--- /dev/null
+++ b/loopstructural/main/workflow_mode.py
@@ -0,0 +1,19 @@
+"""The start choice of the user: how the model is built.
+
+This module does not import QGIS, so the unit tests can run it.
+"""
+
+# Steps 2 (stratigraphy) and 3 (faults) make features and constraints for step 4.
+WORKFLOW_MODE_MAP = 'map'
+# Steps 2 and 3 are hidden. The user gives the constraints in step 4.
+WORKFLOW_MODE_CONSTRAINTS = 'constraints'
+WORKFLOW_MODES = (WORKFLOW_MODE_MAP, WORKFLOW_MODE_CONSTRAINTS)
+DEFAULT_WORKFLOW_MODE = WORKFLOW_MODE_MAP
+
+WORKFLOW_MODE_LABELS = {
+ WORKFLOW_MODE_MAP: "Build from a geological map",
+ WORKFLOW_MODE_CONSTRAINTS: "Interpolate surfaces from constraints",
+}
+
+# The steps that the "constraints" choice hides. The user can show them later.
+CONSTRAINT_MODE_HIDDEN_STEPS = ('stratigraphy', 'faults')
diff --git a/tests/qgis/test_derived_data_state.py b/tests/qgis/test_derived_data_state.py
index 4f4ca42..d4732bb 100644
--- a/tests/qgis/test_derived_data_state.py
+++ b/tests/qgis/test_derived_data_state.py
@@ -267,3 +267,38 @@ def test_the_sources_are_saved_and_loaded(self, data_manager):
data_manager.thickness_sources.clear()
data_manager.update_from_dict(state)
assert data_manager.get_thickness_source(unit.uuid) == derived_data.CALCULATED
+
+
+class TestWorkflowMode:
+ def test_the_default_is_the_map(self, data_manager):
+ assert data_manager.workflow_mode == 'map'
+
+ def test_the_choice_is_saved_and_loaded(self, data_manager):
+ data_manager.set_workflow_mode('constraints')
+ state = json.loads(json.dumps(data_manager.to_dict()))
+ data_manager.set_workflow_mode('map')
+ data_manager.update_from_dict(state)
+ assert data_manager.workflow_mode == 'constraints'
+
+ def test_a_state_file_of_an_older_version_gives_the_map(self, data_manager):
+ data_manager.set_workflow_mode('constraints')
+ state = json.loads(json.dumps(data_manager.to_dict()))
+ del state['workflow_mode']
+ data_manager.update_from_dict(state)
+ assert data_manager.workflow_mode == 'map'
+
+ def test_a_change_calls_the_listeners_one_time(self, data_manager):
+ seen = []
+ data_manager.add_workflow_mode_callback(seen.append)
+ data_manager.set_workflow_mode('constraints')
+ data_manager.set_workflow_mode('constraints')
+ assert seen == ['constraints']
+
+ def test_an_unknown_choice_is_an_error(self, data_manager):
+ with pytest.raises(ValueError):
+ data_manager.set_workflow_mode('other')
+
+ def test_reset_gives_the_map(self, data_manager):
+ data_manager.set_workflow_mode('constraints')
+ data_manager.reset()
+ assert data_manager.workflow_mode == 'map'
diff --git a/tests/qgis/test_detach_feature.py b/tests/qgis/test_detach_feature.py
new file mode 100644
index 0000000..1232f32
--- /dev/null
+++ b/tests/qgis/test_detach_feature.py
@@ -0,0 +1,128 @@
+"""Pytest tests for detaching a feature that the stratigraphic column makes.
+
+A detached feature keeps a copy of its contact and orientation data. Rows that
+the user adds to a generated feature are used together with the data of the
+column.
+"""
+
+import geopandas as gpd
+import pandas as pd
+import pytest
+from LoopStructural import StratigraphicColumn
+
+from loopstructural.main.model_manager import GeologicalModelManager
+
+
+def _contact(x):
+ return pd.DataFrame({'X': [x], 'Y': [0.0], 'Z': [0.0]})
+
+
+@pytest.fixture
+def manager(monkeypatch):
+ manager = GeologicalModelManager()
+ captured = []
+
+ def fake_create_and_add_foliation(name, data=None, **kwargs):
+ captured.append((name, data, kwargs))
+ return object()
+
+ monkeypatch.setattr(manager.model, 'create_and_add_foliation', fake_create_and_add_foliation)
+ monkeypatch.setattr(manager.model, 'add_unconformity', lambda *a, **k: None)
+ manager._captured = captured
+
+ column = StratigraphicColumn()
+ column.clear(basement=False)
+ column.add_unit(name='oldest', thickness=100.0, where='top')
+ column.add_unit(name='youngest', thickness=200.0, where='top')
+ manager.stratigraphic_column = column
+ manager.stratigraphy['oldest']['contact'] = _contact(1.0)
+ manager.stratigraphy['youngest']['contact'] = _contact(2.0)
+ return manager
+
+
+def _group_name(manager):
+ return manager.stratigraphic_column.get_groups()[0].name
+
+
+def _built_x(manager):
+ manager.update_foliation_features()
+ return sorted(manager._captured[-1][1]['X'])
+
+
+class TestDetach:
+ def test_a_group_of_the_column_is_a_generated_feature(self, manager):
+ assert manager.is_generated(_group_name(manager))
+ assert not manager.is_generated('other')
+
+ def test_a_detached_feature_keeps_its_data(self, manager):
+ name = _group_name(manager)
+ assert manager.detach_feature(name)
+ manager.stratigraphy['oldest']['contact'] = _contact(99.0)
+ assert _built_x(manager) == [1.0, 2.0]
+
+ def test_a_feature_that_is_not_detached_follows_the_contacts(self, manager):
+ manager.stratigraphy['oldest']['contact'] = _contact(99.0)
+ assert _built_x(manager) == [2.0, 99.0]
+
+ def test_a_feature_without_data_cannot_be_detached(self, manager):
+ manager.stratigraphy.clear()
+ assert not manager.detach_feature(_group_name(manager))
+
+ def test_a_feature_that_is_not_generated_cannot_be_detached(self, manager):
+ assert not manager.detach_feature('other')
+
+ def test_attach_follows_the_column_again_and_the_model_is_stale(self, manager):
+ name = _group_name(manager)
+ manager.detach_feature(name)
+ manager.stratigraphy['oldest']['contact'] = _contact(99.0)
+ assert manager.attach_feature(name)
+ assert _built_x(manager) == [2.0, 99.0]
+ assert manager._data_dirty
+
+ def test_the_detached_data_is_saved_and_loaded(self, manager):
+ name = _group_name(manager)
+ manager.detach_feature(name)
+ written = manager.detached_to_dict()
+ manager.detached = {}
+ manager.detached_from_dict(written)
+ assert manager.is_detached(name)
+ assert sorted(manager.detached[name]['X']) == [1.0, 2.0]
+
+ def test_reset_forgets_detached_features(self, manager):
+ manager.detach_feature(_group_name(manager))
+ manager.reset()
+ assert manager.detached == {}
+
+
+class TestAddedRows:
+ def _extra(self, x):
+ layer = {
+ 'layer_name': 'values',
+ 'type': 'Value',
+ 'value_field': 'v',
+ 'df': gpd.GeoDataFrame({'v': [5.0]}, geometry=gpd.points_from_xy([x], [0.0])),
+ }
+ return {'data': {'values': layer}, 'use_z_coordinate': False}
+
+ def test_added_rows_are_used_with_the_data_of_the_column(self, manager):
+ name = _group_name(manager)
+ manager.extra_constraints[name] = self._extra(7.0)
+ assert _built_x(manager) == [1.0, 2.0, 7.0]
+
+ def test_added_rows_stay_when_a_feature_is_detached(self, manager):
+ name = _group_name(manager)
+ manager.extra_constraints[name] = self._extra(7.0)
+ manager.detach_feature(name)
+ # the copy has the data of the column only
+ assert sorted(manager.detached[name]['X']) == [1.0, 2.0]
+ assert _built_x(manager) == [1.0, 2.0, 7.0]
+
+ def test_an_inequality_row_gives_the_solver(self, manager):
+ name = _group_name(manager)
+ spec = self._extra(7.0)
+ spec['data']['values'].update(
+ {'type': 'Inequality', 'lower_field': 'v', 'upper_field': 'v'}
+ )
+ manager.extra_constraints[name] = spec
+ manager.update_foliation_features()
+ assert manager._captured[-1][2]['solver'] == 'admm'
diff --git a/tests/qgis/test_manual_foliations.py b/tests/qgis/test_manual_foliations.py
index f699818..8e27583 100644
--- a/tests/qgis/test_manual_foliations.py
+++ b/tests/qgis/test_manual_foliations.py
@@ -136,3 +136,60 @@ def test_loaded_foliation_is_built_again_by_update_model(self, manager):
manager.update_model(notify_observers=False)
assert 's1' in _names(manager)
+
+
+def _points(**fields):
+ points = [Point(x, y, 0.0) for x in (20.0, 50.0, 80.0) for y in (20.0, 50.0, 80.0)]
+ columns = {name: [fn(p) for p in points] for name, fn in fields.items()}
+ return gpd.GeoDataFrame(columns, geometry=points)
+
+
+class TestConstraintTypes:
+ """A feature built from constraint layers only, with no stratigraphic column data."""
+
+ def _rows(self, manager, layer):
+ data, kwargs = manager._foliation_data('f', {'layer': layer}, use_z_coordinate=True)
+ return data, kwargs
+
+ def test_interface_rows_use_the_group_field(self, manager):
+ df = _points(group=lambda p: 'a' if p.x < 50 else 'b')
+ layer = {'type': 'Interface', 'group_field': 'group', 'df': df}
+ data, _ = self._rows(manager, layer)
+ assert sorted(data['interface'].unique()) == [0.0, 1.0]
+
+ def test_weight_goes_to_the_rows(self, manager):
+ layer = _value_layer()
+ layer['weight'] = 0.25
+ data, _ = self._rows(manager, layer)
+ assert (data['w'] == 0.25).all()
+
+ def test_a_constant_z_source_gives_one_z(self, manager):
+ layer = _value_layer()
+ layer.update({'z_source': 'constant', 'z_value': 12.5})
+ data, _ = self._rows(manager, layer)
+ assert (data['Z'] == 12.5).all()
+
+ def test_a_pairwise_inequality_needs_the_admm_solver(self, manager):
+ df = _points(order=lambda p: 1 if p.x < 50 else 2)
+ layer = {'type': 'Pairwise Inequality', 'pair_field': 'order', 'df': df}
+ data, kwargs = self._rows(manager, layer)
+ assert kwargs == {'solver': 'admm'}
+ assert 'pair_id' in data
+
+ def test_a_feature_from_value_and_normal_layers_solves_without_a_column(self, manager):
+ values = _value_layer()
+ normals = {
+ 'layer_name': 'normals',
+ 'type': 'Gradient/Normal',
+ 'vector_kind': 'normal',
+ 'vector_x_field': 'nx',
+ 'vector_y_field': 'ny',
+ 'vector_z_field': 'nz',
+ 'df': _points(nx=lambda p: 1.0, ny=lambda p: 0.0, nz=lambda p: 0.0),
+ }
+ manager.add_foliation(
+ 's1', {'values': values, 'normals': normals}, use_z_coordinate=True
+ )
+ manager.update_all_features(notify_observers=False)
+ result = manager.model['s1'].evaluate_value(np.array([[50.0, 50.0, 0.0]]))
+ assert not np.any(np.isnan(result))
diff --git a/tests/qgis/test_parametric_fault.py b/tests/qgis/test_parametric_fault.py
new file mode 100644
index 0000000..d8fd937
--- /dev/null
+++ b/tests/qgis/test_parametric_fault.py
@@ -0,0 +1,97 @@
+"""Pytest tests for the faults that the user gives by numbers (Add Fault).
+
+The fault has no trace layer. `update_model` builds it together with the other
+faults, so a fault that was added must be built again at each build.
+"""
+
+import json
+
+import numpy as np
+import pytest
+from LoopStructural import StratigraphicColumn
+from LoopStructural.datatypes import BoundingBox
+
+from loopstructural.main.model_manager import GeologicalModelManager
+from loopstructural.toolbelt.preferences import PlgSettingsStructure
+
+
+class _DebugManager:
+ def log(self, *args, **kwargs):
+ pass
+
+
+@pytest.fixture
+def manager(monkeypatch):
+ monkeypatch.setattr(PlgSettingsStructure, 'interpolator_nelements', 200)
+ manager = GeologicalModelManager(debug_manager=_DebugManager())
+ manager.update_bounding_box(BoundingBox(origin=[0, 0, -50], maximum=[100, 100, 50]))
+ manager.stratigraphic_column = StratigraphicColumn()
+ return manager
+
+
+def _spec(**changes):
+ spec = {
+ 'name': 'F1',
+ 'strike': 0.0,
+ 'dip': 90.0,
+ 'pitch': 90.0,
+ 'displacement': 10.0,
+ 'centre': (50.0, 50.0, 0.0),
+ 'major_axis': 100.0,
+ 'intermediate_axis': 100.0,
+ 'minor_axis': 30.0,
+ }
+ spec.update(changes)
+ return spec
+
+
+def _names(manager):
+ return [f.name for f in manager.model.features]
+
+
+class TestParametricFault:
+ def test_the_fault_is_built_by_update_model(self, manager):
+ manager.add_parametric_fault(_spec())
+ manager.update_model(notify_observers=False)
+ assert 'F1' in _names(manager)
+
+ def test_the_fault_surface_is_where_the_fault_is(self, manager):
+ manager.add_parametric_fault(_spec())
+ manager.update_model(notify_observers=False)
+ manager.update_all_features(notify_observers=False)
+ # the fault strikes north through x=50, so the field changes sign across it
+ west = manager.model['F1'][0].evaluate_value(np.array([[30.0, 50.0, 0.0]]))
+ east = manager.model['F1'][0].evaluate_value(np.array([[70.0, 50.0, 0.0]]))
+ assert west[0] * east[0] < 0
+
+ def test_adding_a_fault_makes_the_model_stale(self, manager):
+ manager.add_parametric_fault(_spec())
+ assert manager._data_dirty
+
+ def test_a_name_that_is_used_is_an_error(self, manager):
+ manager.add_parametric_fault(_spec())
+ with pytest.raises(ValueError):
+ manager.add_parametric_fault(_spec())
+
+ def test_a_bad_fault_is_an_error(self, manager):
+ with pytest.raises(ValueError):
+ manager.add_parametric_fault(_spec(dip=0.0))
+
+ def test_a_removed_fault_is_not_built_again(self, manager):
+ manager.add_parametric_fault(_spec())
+ manager.remove_parametric_fault('F1')
+ manager.update_model(notify_observers=False)
+ assert 'F1' not in _names(manager)
+
+ def test_the_faults_are_saved_and_loaded(self, manager):
+ manager.add_parametric_fault(_spec())
+ written = json.loads(json.dumps(manager.parametric_faults_to_dict()))
+ manager.parametric_faults = {}
+ manager.parametric_faults_from_dict(written)
+ manager.update_model(notify_observers=False)
+ assert 'F1' in _names(manager)
+
+ def test_reset_forgets_the_faults(self, manager):
+ manager.add_parametric_fault(_spec())
+ manager.reset()
+ assert manager.parametric_faults == {}
diff --git a/tests/unit/test_build_plan.py b/tests/unit/test_build_plan.py
new file mode 100644
index 0000000..72d6a48
--- /dev/null
+++ b/tests/unit/test_build_plan.py
@@ -0,0 +1,174 @@
+"""Pytest tests for the primary action of the model step and for the choice of
+the derived data that a build calculates again.
+
+The modules do not import QGIS, so the tests use fake managers and run in the
+fast tests/unit/ job.
+"""
+
+from types import SimpleNamespace
+
+import pytest
+
+from loopstructural.gui.modelling.steps import build_plan
+from loopstructural.gui.modelling.steps.status import StepCheck
+from loopstructural.main import derived_data, layer_roles
+from loopstructural.main.derived_data import DerivedData
+from loopstructural.main.derived_refresh import names_to_refresh, thickness_values
+from loopstructural.main.workflow_mode import WORKFLOW_MODE_CONSTRAINTS, WORKFLOW_MODE_MAP
+
+
+class TestPrimaryAction:
+ def test_an_empty_model_is_built(self):
+ action = build_plan.choose_primary_action('empty')
+ assert (action.action, action.text) == (build_plan.ACTION_BUILD, "Build model")
+
+ def test_an_initialized_model_is_solved(self):
+ action = build_plan.choose_primary_action('initialized')
+ assert (action.action, action.text) == (build_plan.ACTION_SOLVE, "Solve model")
+
+ def test_a_solved_model_can_be_solved_again(self):
+ action = build_plan.choose_primary_action('solved')
+ assert action.action == build_plan.ACTION_SOLVE
+ assert action.text == "Solve again"
+
+ def test_a_stale_model_is_rebuilt(self):
+ action = build_plan.choose_primary_action('stale')
+ assert (action.action, action.text) == (build_plan.ACTION_BUILD, "Rebuild model")
+
+ @pytest.mark.parametrize('state', ['initialized', 'solved'])
+ def test_out_of_date_derived_data_makes_a_rebuild(self, state):
+ action = build_plan.choose_primary_action(
+ state, derived_out_of_date=[derived_data.BASAL_CONTACTS]
+ )
+ assert action.action == build_plan.ACTION_BUILD
+ assert 'out of date' in action.tooltip
+
+ def test_a_changed_layer_updates_the_data_and_solves(self):
+ action = build_plan.choose_primary_action('solved', layers_changed=True)
+ assert action.action == build_plan.ACTION_SOLVE
+ assert action.text == "Update data and solve"
+
+ def test_a_blocked_build_is_disabled_with_the_reason(self):
+ action = build_plan.choose_primary_action('solved', blocked_reason="Set the bounding box.")
+ assert action.action == build_plan.ACTION_NONE
+ assert not action.enabled
+ assert action.tooltip == "Set the bounding box."
+
+
+class TestProblems:
+ def test_problems_of_all_steps_in_the_order_of_the_steps(self):
+ problems = build_plan.collect_problems(
+ [
+ ('data', StepCheck(problems=('Bad CRS.',))),
+ ('stratigraphy', StepCheck(todo=('Add units.',))),
+ ('faults', StepCheck(problems=('No faults.',))),
+ ]
+ )
+ assert problems == [('data', 'Bad CRS.'), ('faults', 'No faults.')]
+
+ def test_the_same_problem_shows_one_time(self):
+ problems = build_plan.collect_problems(
+ [
+ ('stratigraphy', StepCheck(problems=('Contacts are out of date.',))),
+ ('model', StepCheck(problems=('Contacts are out of date.',))),
+ ]
+ )
+ assert problems == [('stratigraphy', 'Contacts are out of date.')]
+
+
+class TestModes:
+ ALL = ['data', 'stratigraphy', 'faults', 'model', 'view']
+
+ def test_the_map_choice_shows_all_steps(self):
+ assert build_plan.steps_for_mode(WORKFLOW_MODE_MAP, self.ALL) == self.ALL
+
+ def test_the_constraints_choice_hides_steps_2_and_3(self):
+ shown = build_plan.steps_for_mode(WORKFLOW_MODE_CONSTRAINTS, self.ALL)
+ assert shown == ['data', 'model', 'view']
+
+
+class FakeRoles:
+ def __init__(self):
+ self.values = {layer_roles.GEOLOGY: object(), layer_roles.GEOLOGY_UNIT_FIELD: 'UNITNAME'}
+ self.contacts_source = layer_roles.CONTACTS_FROM_GEOLOGY
+
+
+class FakeDataManager:
+ """The parts of the data manager that `names_to_refresh` reads."""
+
+ def __init__(self):
+ self.units = ['a', 'b']
+ self.layer_roles = FakeRoles()
+ self.inputs = {'unit_order': ['a', 'b']}
+ self.derived = DerivedData()
+ for name in (derived_data.BASAL_CONTACTS, derived_data.THICKNESS):
+ self.derived.register(name, lambda: self.inputs)
+
+ def get_stratigraphic_unit_names(self):
+ return list(self.units)
+
+ def get_layer_role(self, role):
+ return self.layer_roles.values.get(role)
+
+
+@pytest.fixture
+def dm():
+ return FakeDataManager()
+
+
+class TestNamesToRefresh:
+ def test_contacts_that_were_never_calculated_are_calculated(self, dm):
+ assert names_to_refresh(dm) == [derived_data.BASAL_CONTACTS]
+
+ def test_current_contacts_are_not_calculated_again(self, dm):
+ dm.derived.record(derived_data.BASAL_CONTACTS)
+ assert names_to_refresh(dm) == []
+
+ def test_a_reorder_makes_contacts_and_thickness_out_of_date_in_order(self, dm):
+ dm.derived.record(derived_data.BASAL_CONTACTS)
+ dm.derived.record(derived_data.THICKNESS)
+ dm.inputs = {'unit_order': ['b', 'a']}
+ assert names_to_refresh(dm) == [derived_data.BASAL_CONTACTS, derived_data.THICKNESS]
+
+ def test_a_reorder_and_its_undo_calculate_nothing(self, dm):
+ dm.derived.record(derived_data.BASAL_CONTACTS)
+ dm.inputs = {'unit_order': ['b', 'a']}
+ dm.inputs = {'unit_order': ['a', 'b']}
+ assert names_to_refresh(dm) == []
+
+ def test_a_thickness_that_was_never_calculated_can_be_typed(self, dm):
+ dm.derived.record(derived_data.BASAL_CONTACTS)
+ assert derived_data.THICKNESS not in names_to_refresh(dm)
+
+ def test_a_contacts_layer_of_the_user_is_never_calculated(self, dm):
+ dm.layer_roles.contacts_source = layer_roles.CONTACTS_FROM_LAYER
+ assert derived_data.BASAL_CONTACTS not in names_to_refresh(dm)
+
+ def test_no_column_calculates_no_contacts(self, dm):
+ dm.units = []
+ assert names_to_refresh(dm) == []
+
+ def test_no_geology_layer_calculates_no_contacts(self, dm):
+ dm.layer_roles.values[layer_roles.GEOLOGY] = None
+ assert names_to_refresh(dm) == []
+
+
+class TestThicknessValues:
+ def test_the_median_is_used_and_a_unit_with_no_result_is_left_out(self):
+ pd = pytest.importorskip('pandas')
+ table = pd.DataFrame(
+ {
+ 'name': ['a', 'b', 'c'],
+ 'ThicknessMedian': [10.0, -1.0, float('nan')],
+ 'ThicknessMean': [11.0, 12.0, 13.0],
+ }
+ )
+ assert thickness_values(table) == {'a': 10.0}
+
+ def test_the_mean_is_used_without_a_median(self):
+ pd = pytest.importorskip('pandas')
+ table = pd.DataFrame({'name': ['a'], 'ThicknessMean': [5.0]})
+ assert thickness_values(table) == {'a': 5.0}
+
+ def test_a_table_without_thickness_columns_gives_nothing(self):
+ assert thickness_values(SimpleNamespace(columns=['name'])) == {}
diff --git a/tests/unit/test_constraints.py b/tests/unit/test_constraints.py
new file mode 100644
index 0000000..4a28caa
--- /dev/null
+++ b/tests/unit/test_constraints.py
@@ -0,0 +1,173 @@
+"""Pytest tests for the constraint types of a feature.
+
+`loopstructural.main.constraints` does not import QGIS, so the tests run in
+the fast tests/unit/ job.
+"""
+
+import numpy as np
+import pandas as pd
+import pytest
+
+from loopstructural.main import constraints
+
+
+def sampler_of(points):
+ """A sampler that ignores the layer and gives fixed points."""
+ calls = []
+
+ def sampler(df, dem, use_z):
+ calls.append(use_z)
+ return points.copy()
+
+ sampler.calls = calls
+ return sampler
+
+
+@pytest.fixture
+def points():
+ return pd.DataFrame(
+ {
+ 'X': [0.0, 1.0, 2.0, 3.0],
+ 'Y': [0.0, 1.0, 2.0, 3.0],
+ 'Z': [10.0, 11.0, 12.0, 13.0],
+ 'feature_id': [0, 0, 1, 1],
+ 'gx': [0.0, 0.0, 1.0, np.nan],
+ 'gy': [0.0, 0.0, 0.0, 1.0],
+ 'gz': [1.0, 0.0, 0.0, 1.0],
+ 'group': ['a', 'b', 'a', None],
+ 'order': [1, 1, 2, 2],
+ }
+ )
+
+
+class TestTypes:
+ def test_every_type_has_a_description_and_a_field_rule(self):
+ for layer_type in constraints.CONSTRAINT_TYPES:
+ assert constraints.DESCRIPTIONS[layer_type]
+ assert layer_type in constraints.REQUIRED_FIELDS
+
+ def test_the_six_types_of_the_plan_are_there(self):
+ for name in (
+ 'Value',
+ 'Interface',
+ 'Gradient/Normal',
+ 'Tangent',
+ 'Inequality',
+ 'Pairwise Inequality',
+ ):
+ assert name in constraints.CONSTRAINT_TYPES
+
+ def test_missing_fields(self):
+ row = {'type': constraints.GRADIENT_NORMAL, 'vector_x_field': 'gx'}
+ assert constraints.missing_fields(row) == ['vector_y_field', 'vector_z_field']
+ assert constraints.missing_fields({'type': constraints.INTERFACE}) == []
+
+ def test_only_inequality_types_need_the_admm_solver(self):
+ assert constraints.solver_for(constraints.INEQUALITY) == 'admm'
+ assert constraints.solver_for(constraints.PAIRWISE_INEQUALITY) == 'admm'
+ assert constraints.solver_for(constraints.VALUE) is None
+
+
+class TestZSource:
+ def test_a_row_of_an_older_version_uses_the_default(self, points):
+ sampler = sampler_of(points)
+ constraints.sample_layer(sampler, {'df': None}, None, default_use_z=True)
+ assert sampler.calls == [True]
+
+ @pytest.mark.parametrize(
+ 'source, expected', [('layer', True), ('dem', False), ('constant', False)]
+ )
+ def test_the_source_decides_if_the_layer_z_is_used(self, points, source, expected):
+ sampler = sampler_of(points)
+ constraints.sample_layer(sampler, {'df': None, 'z_source': source}, None, True)
+ assert sampler.calls == [expected]
+
+ def test_a_constant_replaces_all_z(self, points):
+ rows = constraints.sample_layer(
+ sampler_of(points), {'df': None, 'z_source': 'constant', 'z_value': -50}, None
+ )
+ assert (rows['Z'] == -50.0).all()
+
+
+class TestWeight:
+ def test_no_weight_adds_no_column(self, points):
+ assert 'w' not in constraints.add_weight(points[['X', 'Y', 'Z']], points, {})
+
+ def test_the_weight_is_added(self, points):
+ rows = constraints.add_weight(points[['X', 'Y', 'Z']], points, {'weight': 0.5})
+ assert (rows['w'] == 0.5).all()
+
+ def test_a_weight_in_the_rows_is_kept(self, points):
+ rows = points[['X', 'Y', 'Z']].assign(w=0.1)
+ assert (constraints.add_weight(rows, points, {'weight': 5})['w'] == 0.1).all()
+
+
+class TestInterface:
+ def test_each_feature_is_a_surface_without_a_group_field(self, points):
+ rows, offset = constraints.interface_rows(points, {}, 'f')
+ assert list(rows['interface']) == [0, 0, 1, 1]
+ assert offset == 2
+
+ def test_the_offset_keeps_the_surfaces_of_two_rows_apart(self, points):
+ rows, offset = constraints.interface_rows(points, {}, 'f', offset=5)
+ assert list(rows['interface']) == [5, 5, 6, 6]
+ assert offset == 7
+
+ def test_a_group_field_gives_the_surfaces_and_drops_empty_values(self, points):
+ rows, offset = constraints.interface_rows(points, {'group_field': 'group'}, 'f')
+ assert list(rows['interface']) == [0, 1, 0]
+ assert offset == 2
+ assert (rows['feature_name'] == 'f').all()
+
+
+class TestVectors:
+ def test_a_gradient_gives_the_g_columns_and_drops_bad_rows(self, points):
+ row = {
+ 'type': constraints.GRADIENT_NORMAL,
+ 'vector_x_field': 'gx',
+ 'vector_y_field': 'gy',
+ 'vector_z_field': 'gz',
+ }
+ rows = constraints.constraint_rows(points, row, 'f')
+ # row 1 is a zero vector and row 3 has no x
+ assert list(rows.index) == [0, 2]
+ assert {'gx', 'gy', 'gz'} <= set(rows.columns)
+
+ def test_a_normal_gives_the_n_columns(self, points):
+ row = {
+ 'type': constraints.GRADIENT_NORMAL,
+ 'vector_kind': constraints.KIND_NORMAL,
+ 'vector_x_field': 'gx',
+ 'vector_y_field': 'gy',
+ 'vector_z_field': 'gz',
+ 'weight': 2.0,
+ }
+ rows = constraints.constraint_rows(points, row, 'f')
+ assert {'nx', 'ny', 'nz'} <= set(rows.columns)
+ assert (rows['w'] == 2.0).all()
+
+ def test_a_tangent_gives_the_t_columns(self, points):
+ row = {
+ 'type': constraints.TANGENT,
+ 'vector_x_field': 'gx',
+ 'vector_y_field': 'gy',
+ 'vector_z_field': 'gz',
+ }
+ assert {'tx', 'ty', 'tz'} <= set(constraints.constraint_rows(points, row, 'f').columns)
+
+
+class TestPairwise:
+ def test_the_group_number_gives_the_pair_id(self, points):
+ row = {'type': constraints.PAIRWISE_INEQUALITY, 'pair_field': 'order'}
+ rows = constraints.constraint_rows(points, row, 'f')
+ assert list(rows['pair_id']) == [1, 1, 2, 2]
+
+ def test_a_row_without_a_number_is_dropped(self, points):
+ points.loc[0, 'order'] = np.nan
+ row = {'type': constraints.PAIRWISE_INEQUALITY, 'pair_field': 'order'}
+ assert len(constraints.constraint_rows(points, row, 'f')) == 3
+
+
+def test_an_unknown_type_is_an_error(points):
+ with pytest.raises(ValueError):
+ constraints.constraint_rows(points, {'type': 'Other'}, 'f')
diff --git a/tests/unit/test_parametric_fault.py b/tests/unit/test_parametric_fault.py
new file mode 100644
index 0000000..b6699c9
--- /dev/null
+++ b/tests/unit/test_parametric_fault.py
@@ -0,0 +1,113 @@
+"""Pytest tests for the faults that the user gives by numbers.
+
+`loopstructural.main.parametric_fault` does not import QGIS, so the tests run
+in the fast tests/unit/ job.
+"""
+
+import json
+
+import numpy as np
+import pytest
+
+from loopstructural.main import parametric_fault
+
+
+def good_spec(**changes):
+ spec = {
+ 'name': 'F1',
+ 'strike': 0.0,
+ 'dip': 60.0,
+ 'pitch': 90.0,
+ 'displacement': 50.0,
+ 'centre': (5.0, 6.0, 7.0),
+ 'major_axis': 100.0,
+ 'intermediate_axis': 80.0,
+ 'minor_axis': 20.0,
+ }
+ spec.update(changes)
+ return spec
+
+
+class TestPlaneVectors:
+ def test_a_vertical_plane_that_strikes_north_has_an_east_normal(self):
+ normal, _ = parametric_fault.plane_vectors(0, 90, 0)
+ assert np.allclose(normal, [1, 0, 0], atol=1e-9)
+
+ def test_a_flat_plane_has_an_up_normal(self):
+ normal, _ = parametric_fault.plane_vectors(123, 0, 0)
+ assert np.allclose(normal, [0, 0, 1], atol=1e-9)
+
+ def test_the_plane_dips_to_the_right_of_the_strike(self):
+ # strike east, so the dip direction is south: the normal tilts to the south
+ normal, _ = parametric_fault.plane_vectors(90, 45, 0)
+ assert normal[1] < 0 < normal[2]
+
+ @pytest.mark.parametrize('strike', [0, 37, 90, 210])
+ @pytest.mark.parametrize('dip', [10, 45, 90])
+ @pytest.mark.parametrize('pitch', [0, 30, 90, -45])
+ def test_the_slip_is_a_unit_vector_in_the_plane(self, strike, dip, pitch):
+ normal, slip = parametric_fault.plane_vectors(strike, dip, pitch)
+ assert np.isclose(np.linalg.norm(normal), 1)
+ assert np.isclose(np.linalg.norm(slip), 1)
+ assert np.isclose(normal @ slip, 0, atol=1e-9)
+
+ def test_pitch_0_is_along_the_strike_and_pitch_90_is_down_the_dip(self):
+ _, along = parametric_fault.plane_vectors(0, 90, 0)
+ _, down = parametric_fault.plane_vectors(0, 90, 90)
+ assert np.allclose(along, [0, 1, 0], atol=1e-9)
+ assert np.allclose(down, [0, 0, -1], atol=1e-9)
+
+
+class TestProblems:
+ def test_a_good_fault_has_no_problem(self):
+ assert parametric_fault.problems(good_spec()) == []
+
+ def test_the_name_is_needed_and_must_be_new(self):
+ assert parametric_fault.problems(good_spec(name=' '))
+ assert 'already used' in parametric_fault.problems(good_spec(), ['F1'])[0]
+
+ @pytest.mark.parametrize('dip', [0, -5, 91])
+ def test_the_dip_must_be_in_range(self, dip):
+ assert parametric_fault.problems(good_spec(dip=dip))
+
+ @pytest.mark.parametrize('axis', ['major_axis', 'intermediate_axis', 'minor_axis'])
+ def test_a_size_must_be_above_zero(self, axis):
+ assert parametric_fault.problems(good_spec(**{axis: 0}))
+
+ def test_the_centre_needs_three_numbers(self):
+ assert parametric_fault.problems(good_spec(centre=(1.0, 2.0)))
+ assert parametric_fault.problems(good_spec(centre=(1.0, 2.0, float('nan'))))
+
+
+class TestLoopStructuralArguments:
+ def test_the_frame_data_is_the_centre_on_the_fault_surface(self):
+ data = parametric_fault.frame_data(good_spec())
+ row = data.iloc[0]
+ assert (row['X'], row['Y'], row['Z']) == (5.0, 6.0, 7.0)
+ assert row['val'] == 0 and row['coord'] == 0 and row['feature_name'] == 'F1'
+
+ def test_the_arguments(self):
+ arguments = parametric_fault.fault_arguments(good_spec())
+ assert arguments['displacement'] == 50.0
+ assert arguments['major_axis'] == 100.0
+ assert arguments['intermediate_axis'] == 80.0
+ assert arguments['minor_axis'] == 20.0
+ assert list(arguments['fault_center']) == [5.0, 6.0, 7.0]
+ assert np.isclose(arguments['fault_normal_vector'] @ arguments['fault_slip_vector'], 0)
+
+
+class TestSaving:
+ def test_a_spec_survives_json(self):
+ spec = good_spec(name=' F1 ', centre=(np.float64(5.0), 6, 7))
+ written = json.loads(json.dumps(parametric_fault.clean_spec(spec)))
+ loaded = parametric_fault.spec_from_dict(written)
+ assert loaded['name'] == 'F1'
+ assert loaded['centre'] == (5.0, 6.0, 7.0)
+ assert loaded['dip'] == 60.0
+
+ def test_the_default_spec_is_in_the_middle_of_the_area(self):
+ spec = parametric_fault.default_spec([0, 0, -100], [1000, 400, 100])
+ assert spec['centre'] == (500.0, 200.0, 0.0)
+ assert spec['major_axis'] == 1000.0
+ assert spec['minor_axis'] == 500.0
+ assert parametric_fault.problems({**spec, 'name': 'F'}) == []
diff --git a/tests/unit/test_preview.py b/tests/unit/test_preview.py
new file mode 100644
index 0000000..4665200
--- /dev/null
+++ b/tests/unit/test_preview.py
@@ -0,0 +1,79 @@
+"""Pytest tests for the map preview of a feature (grid, levels, isolines).
+
+`loopstructural.main.preview` does not import QGIS, so the tests run in the
+fast tests/unit/ job.
+"""
+
+import numpy as np
+import pytest
+
+pytest.importorskip('contourpy')
+
+from loopstructural.main import preview
+
+
+class TestGrid:
+ def test_the_points_go_row_by_row(self):
+ x, y, points = preview.grid_points([0, 10, -5], [4, 20, 5], 3)
+ assert list(x) == [0, 2, 4]
+ assert list(y) == [10, 15, 20]
+ assert points.shape == (9, 3)
+ assert list(points[:3, 1]) == [10, 10, 10]
+ assert list(points[:3, 0]) == [0, 2, 4]
+
+ def test_z_is_the_dem(self):
+ _, _, points = preview.grid_points([0, 0], [1, 1], 2, lambda x, y: 100 + x + 2 * y)
+ assert sorted(points[:, 2]) == [100, 101, 102, 103]
+
+ def test_without_a_dem_z_is_zero(self):
+ _, _, points = preview.grid_points([0, 0], [1, 1], 2)
+ assert (points[:, 2] == 0).all()
+
+ def test_a_grid_needs_two_points_on_each_side(self):
+ with pytest.raises(ValueError):
+ preview.grid_points([0, 0], [1, 1], 1)
+
+
+class TestLevels:
+ def test_the_levels_are_inside_the_range(self):
+ levels = preview.default_levels([0.0, 10.0, np.nan], 4)
+ assert len(levels) == 4
+ assert levels.min() > 0 and levels.max() < 10
+ assert np.allclose(np.diff(levels), np.diff(levels)[0])
+
+ @pytest.mark.parametrize('values', [[], [np.nan], [5.0, 5.0]])
+ def test_a_field_without_a_range_has_no_levels(self, values):
+ assert len(preview.default_levels(values)) == 0
+
+
+class TestIsolines:
+ def setup_method(self):
+ self.x = np.linspace(0, 10, 11)
+ self.y = np.linspace(0, 10, 11)
+ # a field that grows with X: the lines are straight and run along Y
+ self.values = np.tile(self.x, (len(self.y), 1))
+
+ def test_a_line_is_at_the_level(self):
+ lines = preview.isolines(self.x, self.y, self.values, [2.5])
+ assert len(lines) == 1
+ level, coordinates = lines[0]
+ assert level == 2.5
+ assert np.allclose(coordinates[:, 0], 2.5)
+ assert coordinates[:, 1].min() == 0 and coordinates[:, 1].max() == 10
+
+ def test_a_level_outside_the_range_has_no_line(self):
+ assert preview.isolines(self.x, self.y, self.values, [50.0]) == []
+
+ def test_a_gap_stops_the_line(self):
+ values = self.values.copy()
+ values[:, 2:4] = np.nan
+ assert preview.isolines(self.x, self.y, values, [2.5]) == []
+
+ def test_the_shape_must_fit_the_grid(self):
+ with pytest.raises(ValueError):
+ preview.isolines(self.x, self.y, self.values.T[:, :5], [1.0])
+
+ def test_the_points_of_the_grid_make_the_same_lines(self):
+ x, y, points = preview.grid_points([0, 0], [10, 10], 11)
+ values = points[:, 0].reshape(len(y), len(x))
+ assert len(preview.isolines(x, y, values, preview.default_levels(values, 3))) == 3