From 13b8e4a2d55d2a26dd21c54083a771437ec47949 Mon Sep 17 00:00:00 2001 From: Sanjeev Bashyal Date: Fri, 11 Sep 2026 14:11:50 +0200 Subject: [PATCH 1/6] added support for nml-tools gui using qtpy --- README.md | 51 +++ pyproject.toml | 5 + src/nml_tools/cli.py | 28 ++ src/nml_tools/gui/__init__.py | 27 ++ src/nml_tools/gui/app.py | 544 +++++++++++++++++++++++++++ src/nml_tools/gui/arrays.py | 257 +++++++++++++ src/nml_tools/gui/fields.py | 478 ++++++++++++++++++++++++ src/nml_tools/gui/model.py | 674 ++++++++++++++++++++++++++++++++++ tests/test_cli_gui.py | 36 ++ tests/test_gui_arrays.py | 141 +++++++ tests/test_gui_model.py | 266 ++++++++++++++ tests/test_gui_widgets.py | 183 +++++++++ 12 files changed, 2690 insertions(+) create mode 100644 src/nml_tools/gui/__init__.py create mode 100644 src/nml_tools/gui/app.py create mode 100644 src/nml_tools/gui/arrays.py create mode 100644 src/nml_tools/gui/fields.py create mode 100644 src/nml_tools/gui/model.py create mode 100644 tests/test_cli_gui.py create mode 100644 tests/test_gui_arrays.py create mode 100644 tests/test_gui_model.py create mode 100644 tests/test_gui_widgets.py diff --git a/README.md b/README.md index 0ac4e6f..5bc70b6 100644 --- a/README.md +++ b/README.md @@ -861,6 +861,57 @@ Character substring assignment, complex schema values, nested derived values, component arrays, nondefault lower bounds, and user-defined formatted I/O are currently explicit capability boundaries. +## GUI + +Install the current checkout with its optional GUI dependencies and a Qt binding +(Python 3.9 or newer): + +```bash +python -m pip install '.[gui]' PyQt5 +nml-tools gui -i /path/to/schemas -o /path/to/project +``` + +QtPy also permits other supported Qt bindings. Omit `-i` to use the current +directory, and `-o` to save alongside `nml-config.toml`. + +The editor loads each profile's `default_file` from the output directory and +saves directly to that namelist file. Missing values use schema defaults, then +examples or type-specific suggestions. Invalid input is reported. No JSON +configuration or JSON-to-namelist conversion is involved. Selected groups are +rewritten on save; other groups and their comments remain intact. + +Each profile has a tab with a namelist list, editable pages, navigation, reset, +cancel, and save actions. The Config tab applies runtime dimensions with Run; +the `+` tab can import a `.nml` file or create another profile. One-element +one-dimensional arrays use inline scalar or derived-object editors while +remaining arrays in the saved namelist. Resizable arrays retain a resize action. + +Applications can launch a subset of the configured profiles: + +```python +from nml_tools.gui import launch_gui + +launch_gui( + schemas_dir="/path/to/schemas", + output_dir="/path/to/project", + file_profiles={"main": ["mainconfig", "time_periods"], "parameter": []}, + initial_values={"main": {"mainconfig": {"nDomains": 2}}}, + initial_dimensions={"max_domains": 2}, +) +``` + +`file_profiles` is the third argument. `None` or `{}` selects all configured +profiles; an empty list selects every namelist in that profile. Names are +checked case-insensitively and pages retain TOML order. `initial_values` is an +in-memory dictionary of profile/namelist/field values applied over existing +input. Use keyword arguments for values and dimensions when migrating callers +of the previous GUI API. + +Named runtime dimensions come from TOML, `initial_dimensions`, or the Config +controls; partial namelist assignments cannot reliably recover them. Existing +Qt applications reuse their QApplication; importing `nml_tools.gui` alone +does not import Qt. + ## Error handling Generated type-bound procedures return integer status codes and accept an diff --git a/pyproject.toml b/pyproject.toml index b20345f..c73c344 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,6 +20,11 @@ dependencies = [ ] [project.optional-dependencies] +gui = [ + "guidata; python_version >= '3.9'", + "numpy; python_version >= '3.9'", + "qtpy; python_version >= '3.9'", +] dev = [ "numpy>=1.24,<3", "pytest-cov", diff --git a/src/nml_tools/cli.py b/src/nml_tools/cli.py index 97861e6..df4d126 100644 --- a/src/nml_tools/cli.py +++ b/src/nml_tools/cli.py @@ -1279,6 +1279,34 @@ def cli(verbose: int, quiet: int) -> None: _configure_logging(verbose, quiet) +@cli.command("gui", context_settings=_CONTEXT_SETTINGS) +@click.option( + "--input-path", "-i", + type=click.Path(exists=True, file_okay=False, path_type=Path), + help="Directory containing nml-config.toml and schemas (default: current directory).", +) +@click.option( + "--output-path", "-o", + type=click.Path(file_okay=False, path_type=Path), + help="Directory for namelist files (default: input path).", +) +def gui(input_path: Path | None, output_path: Path | None) -> None: + """Edit and save namelist files using schema-driven Qt forms.""" + try: + from .gui import launch_gui + + exit_code = launch_gui(input_path, output_path) + except ImportError as exc: + raise click.ClickException( + "GUI dependencies are unavailable; install 'nml-tools[gui]' " + f"and a Qt binding: {exc}" + ) from exc + except (OSError, RuntimeError, ValueError) as exc: + raise click.ClickException(f"failed to start GUI: {exc}") from exc + if exit_code: + raise Exit(exit_code) + + @cli.command("generate", context_settings=_CONTEXT_SETTINGS) @click.option( "--config", diff --git a/src/nml_tools/gui/__init__.py b/src/nml_tools/gui/__init__.py new file mode 100644 index 0000000..3282775 --- /dev/null +++ b/src/nml_tools/gui/__init__.py @@ -0,0 +1,27 @@ +"""Optional Qt namelist editor; Qt is imported only when launching the GUI.""" + +from __future__ import annotations + +from collections.abc import Mapping +from pathlib import Path +from typing import Any + +__all__ = ["launch_gui"] + + +def launch_gui( + schemas_dir: Path | str | None = None, + output_dir: Path | str | None = None, + file_profiles: Mapping[str, list[str]] | None = None, + initial_values: Mapping[str, Any] | None = None, + initial_dimensions: Mapping[str, int] | None = None, +) -> int: + """Edit selected profiles; empty lists select all namelists in that profile. + + None or an empty mapping selects all configured profiles. Values are nested + by profile, namelist, and field, and override existing namelist input. + Dimensions override TOML defaults. Output defaults to schemas_dir. + """ + from .app import launch_gui as _launch_gui + + return _launch_gui(schemas_dir, output_dir, file_profiles, initial_values, initial_dimensions) diff --git a/src/nml_tools/gui/app.py b/src/nml_tools/gui/app.py new file mode 100644 index 0000000..60b47cb --- /dev/null +++ b/src/nml_tools/gui/app.py @@ -0,0 +1,544 @@ +"""Tabbed Qt editor for namelist files.""" + +from __future__ import annotations + +import copy +import sys +from collections.abc import Mapping +from pathlib import Path +from typing import Any + +from qtpy.QtCore import Qt +from qtpy.QtWidgets import ( + QAbstractItemView, + QApplication, + QComboBox, + QDialog, + QFileDialog, + QFormLayout, + QGroupBox, + QHBoxLayout, + QLabel, + QLineEdit, + QListWidget, + QListWidgetItem, + QMessageBox, + QPushButton, + QScrollArea, + QSpinBox, + QStackedWidget, + QTabBar, + QTabWidget, + QVBoxLayout, + QWidget, +) + +from .fields import NamelistForm, _exec +from .model import ( + GuiProfile, + GuiProject, + _normalize_dimensions, + _normalize_profile_values, + create_virtual_project, + import_profile, + load_profile, + load_project, + overlay_values, + save_profiles, +) + + +class ProfileTab(QWidget): + """Ordered namelist pages for one file profile.""" + + def __init__( + self, + project: GuiProject, + profile: GuiProfile, + values: Mapping[str, Any], + dimensions: Mapping[str, int], + parent: QWidget | None = None, + *, + fit_arrays: bool = False, + ): + super().__init__(parent) + self.profile = profile + sizes = {**project.constants, **dimensions} + + root = QVBoxLayout(self) + if profile.description: + description = QLabel(profile.description, self) + description.setWordWrap(True) + root.addWidget(description) + + pages = QHBoxLayout() + self.selector = QListWidget(self) + self.selector.setMaximumWidth(240) + self.stack = QStackedWidget(self) + self.forms: dict[str, NamelistForm] = {} + for page in profile.pages: + item = QListWidgetItem(page.name) + item.setData(Qt.ItemDataRole.UserRole, page.key) + title = page.schema.get("title") + if isinstance(title, str): + item.setToolTip(title) + self.selector.addItem(item) + form = NamelistForm( + page.schema, + values.get(page.name), + sizes, + fit_arrays=fit_arrays, + ) + scroll = QScrollArea(self) + scroll.setWidgetResizable(True) + scroll.setWidget(form) + self.stack.addWidget(scroll) + self.forms[page.name] = form + self.selector.currentRowChanged.connect(self.stack.setCurrentIndex) + pages.addWidget(self.selector) + pages.addWidget(self.stack, 1) + root.addLayout(pages, 1) + + buttons = QHBoxLayout() + self.back = QPushButton("Back", self) + self.next = QPushButton("Next", self) + self.restore = QPushButton("Restore defaults", self) + self.cancel = QPushButton("Cancel", self) + self.save = QPushButton("Save", self) + self.back.clicked.connect( + lambda: self.selector.setCurrentRow(self.selector.currentRow() - 1) + ) + self.next.clicked.connect( + lambda: self.selector.setCurrentRow(self.selector.currentRow() + 1) + ) + self.restore.clicked.connect(self.restore_page) + buttons.addWidget(self.back) + buttons.addWidget(self.next) + buttons.addStretch(1) + buttons.addWidget(self.restore) + buttons.addWidget(self.cancel) + buttons.addWidget(self.save) + root.addLayout(buttons) + + self.selector.currentRowChanged.connect(self._update_navigation) + if profile.pages: + self.selector.setCurrentRow(0) + else: + self._update_navigation(-1) + + def values(self) -> dict[str, Any]: + """Return the values of every namelist page.""" + return {page.name: self.forms[page.name].values() for page in self.profile.pages} + + def restore_page(self) -> None: + """Restore defaults on the currently selected page.""" + index = self.selector.currentRow() + if index >= 0: + page = self.profile.pages[index] + self.forms[page.name].reset() + + def restore_all(self) -> None: + """Restore defaults on every page.""" + for form in self.forms.values(): + form.reset() + + def _update_navigation(self, index: int) -> None: + self.back.setEnabled(index > 0) + self.next.setEnabled(0 <= index < len(self.profile.pages) - 1) + self.restore.setEnabled(index >= 0) + + +class ProfileConfigTab(QWidget): + """Inputs used to activate or create file-profile tabs.""" + + def __init__( + self, + project: GuiProject, + dimensions: Mapping[str, int], + parent: QWidget, + *, + builder: bool, + primary: bool = False, + ): + super().__init__(parent) + self.builder = builder + self.primary = primary + self.source_path: Path | None = None + self.source_combo = QComboBox(self) + self.browse = QPushButton("Browse…", self) + self.dimension_boxes: dict[str, QSpinBox] = {} + self.available_schemas: QListWidget | None = None + self.selected_schemas: QListWidget | None = None + self.profile_name: QLineEdit | None = None + self.default_filename: QLineEdit | None = None + + layout = QVBoxLayout(self) + source_layout = QHBoxLayout() + source_layout.addWidget(QLabel("Load configuration", self)) + source_layout.addWidget(self.source_combo, 1) + source_layout.addWidget(self.browse) + layout.addLayout(source_layout) + + dimensions_group = QGroupBox("Runtime dimensions", self) + dimensions_layout = QFormLayout(dimensions_group) + for name, default in project.default_dimensions.items(): + box = QSpinBox(dimensions_group) + box.setRange(1, 2_147_483_647) + box.setValue(dimensions.get(name, default)) + dimensions_layout.addRow(name, box) + self.dimension_boxes[name] = box + if self.dimension_boxes: + layout.addWidget(dimensions_group) + else: + dimensions_group.hide() + + if builder: + self._build_profile_controls(project, layout) + layout.addStretch(1) + run_row = QHBoxLayout() + run_row.addStretch(1) + self.run = QPushButton("Run", self) + run_row.addWidget(self.run) + layout.addLayout(run_row) + + def _build_profile_controls(self, project: GuiProject, layout: QVBoxLayout) -> None: + group = QGroupBox("File profile", self) + group_layout = QVBoxLayout(group) + lists = QHBoxLayout() + available = QListWidget(group) + selected = QListWidget(group) + self.available_schemas = available + self.selected_schemas = selected + selection_mode = QAbstractItemView.ExtendedSelection + available.setSelectionMode(selection_mode) + selected.setSelectionMode(selection_mode) + for page in project.namelists: + available.addItem(ConfigurationDialog._schema_item(page.name, page.key)) + + transfers = QVBoxLayout() + transfers.addStretch(1) + for label, source, target, move_all in ( + (">", available, selected, False), + (">>", available, selected, True), + ("<", selected, available, False), + ("<<", selected, available, True), + ): + button = QPushButton(label, group) + handler = ( + ConfigurationDialog._move_all if move_all else ConfigurationDialog._move_selected + ) + button.clicked.connect(lambda _checked=False, s=source, t=target, h=handler: h(s, t)) + transfers.addWidget(button) + transfers.addStretch(1) + lists.addWidget(available, 1) + lists.addLayout(transfers) + lists.addWidget(selected, 1) + group_layout.addLayout(lists) + + metadata = QFormLayout() + self.profile_name = QLineEdit(group) + self.default_filename = QLineEdit(group) + metadata.addRow("Profile name", self.profile_name) + metadata.addRow("Default file name", self.default_filename) + group_layout.addLayout(metadata) + layout.addWidget(group, 1) + + def dimensions(self) -> dict[str, int]: + return {name: box.value() for name, box in self.dimension_boxes.items()} + + def schema_keys(self) -> list[str]: + if self.selected_schemas is None: + return [] + result: list[str] = [] + for index in range(self.selected_schemas.count()): + item = self.selected_schemas.item(index) + if item is not None: + result.append(str(item.data(Qt.ItemDataRole.UserRole))) + return result + + +class ConfigurationDialog(QDialog): + """Edit configured profiles and optional imported or user-created namelists.""" + + def __init__( + self, + project: GuiProject, + parent: QWidget | None = None, + initial_values: Mapping[str, Any] | None = None, + initial_dimensions: Mapping[str, int] | None = None, + ): + super().__init__(parent) + self.project = project + self.dimensions = _normalize_dimensions( + {} if initial_dimensions is None else initial_dimensions, project + ) + self.initial_values: dict[str, Any] = {} + if initial_values is not None: + if not isinstance(initial_values, Mapping): + raise ValueError("initial_values must map profile names to namelist values") + for name, values in initial_values.items(): + if not isinstance(name, str): + raise ValueError("initial value profile names must be strings") + try: + profile = project.profile(name) + except KeyError as exc: + raise ValueError(f"unknown initial value profile '{name}'") from exc + if profile.key in self.initial_values: + raise ValueError(f"duplicate initial value profile '{name}'") + self.initial_values[profile.key] = _normalize_profile_values( + values, profile, {**project.constants, **self.dimensions} + ) + self.editors: dict[Path, ProfileTab] = {} + self.config_tabs: set[ProfileConfigTab] = set() + self.setWindowTitle("Namelist configuration") + self.resize(1000, 700) + root = QVBoxLayout(self) + self.tabs = QTabWidget(self) + self.tabs.setTabsClosable(True) + root.addWidget(self.tabs, 1) + self.plus_tab = QWidget(self.tabs) + self.tabs.addTab(self.plus_tab, "+") + self.config_tab = self._add_config(primary=True) + for side in (QTabBar.LeftSide, QTabBar.RightSide): + self.tabs.tabBar().setTabButton(self.tabs.indexOf(self.plus_tab), side, None) + self.tabs.currentChanged.connect(self._tab_changed) + self.tabs.tabCloseRequested.connect(self._close_tab) + actions = QHBoxLayout() + actions.addStretch(1) + for label, callback in ( + ("Restore all", self._restore_all), + ("Save all", self._save_all), + ("Close", self.accept), + ): + button = QPushButton(label, self) + button.clicked.connect(callback) + actions.addWidget(button) + root.addLayout(actions) + if project.profiles: + self._run_configuration(self.config_tab) + + @staticmethod + def _schema_item(name: str, key: str) -> QListWidgetItem: + item = QListWidgetItem(name) + item.setData(Qt.ItemDataRole.UserRole, key) + return item + + @staticmethod + def _move_selected(source: QListWidget, target: QListWidget) -> None: + for item in source.selectedItems(): + target.addItem(source.takeItem(source.row(item))) + + @staticmethod + def _move_all(source: QListWidget, target: QListWidget) -> None: + while source.count(): + target.addItem(source.takeItem(0)) + + def _add_config(self, *, primary: bool = False) -> ProfileConfigTab: + tab = ProfileConfigTab( + self.project, + self.dimensions, + self, + builder=not primary or not self.project.profiles, + primary=primary, + ) + tab.source_combo.addItem("Configured profiles" if not tab.builder else "New file profile") + for path in sorted(self.project.output_root.glob("*.nml")): + tab.source_combo.addItem(path.name, str(path)) + tab.browse.clicked.connect(lambda: self._browse(tab)) + tab.source_combo.currentIndexChanged.connect(lambda: self._select_source(tab)) + tab.run.clicked.connect(lambda: self._run_configuration(tab)) + self.config_tabs.add(tab) + self.tabs.insertTab(self.tabs.indexOf(self.plus_tab), tab, "Config") + self.tabs.setCurrentWidget(tab) + return tab + + def _tab_changed(self, index: int) -> None: + if self.tabs.widget(index) is self.plus_tab: + self.tabs.blockSignals(True) + try: + self._add_config() + finally: + self.tabs.blockSignals(False) + + def _close_tab(self, index: int) -> None: + widget = self.tabs.widget(index) + if widget is self.plus_tab: + return + try: + dirty = isinstance(widget, ProfileTab) and widget.values() != widget.saved_values + except ValueError: + dirty = True + if dirty: + if ( + QMessageBox.question(self, "Unsaved changes", "Discard changes in this profile?") + != QMessageBox.Yes + ): + return + self.editors = {path: tab for path, tab in self.editors.items() if tab is not widget} + self.config_tabs.discard(widget) + self.tabs.removeTab(index) + widget.deleteLater() + + def _browse(self, tab: ProfileConfigTab) -> None: + name, _ = QFileDialog.getOpenFileName( + self, "Load namelist", str(self.project.output_root), "Namelist files (*.nml)" + ) + if name: + index = tab.source_combo.findData(name) + if index < 0: + tab.source_combo.addItem(Path(name).name, name) + index = tab.source_combo.count() - 1 + tab.source_combo.setCurrentIndex(index) + + def _select_source(self, tab: ProfileConfigTab) -> None: + name = tab.source_combo.currentData() + tab.source_path = Path(name) if name else None + if not tab.builder or not name: + return + try: + profile, _ = import_profile(self.project, Path(name), tab.dimensions()) + except (OSError, ValueError) as exc: + QMessageBox.critical(self, "Invalid namelist", str(exc)) + return + tab.profile_name.setText(profile.name) + tab.default_filename.setText(profile.default_file) + tab.available_schemas.clear() + tab.selected_schemas.clear() + selected = {page.key for page in profile.pages} + for page in self.project.namelists: + target = tab.selected_schemas if page.key in selected else tab.available_schemas + target.addItem(self._schema_item(page.name, page.key)) + + def _run_configuration(self, config: ProfileConfigTab) -> None: + prepared: list[ProfileTab] = [] + try: + dimensions = _normalize_dimensions(config.dimensions(), self.project) + imported: dict[str, Any] | None = None + if config.source_path is not None: + profile, imported = import_profile(self.project, config.source_path, dimensions) + profiles = (profile,) + if config.builder: + profiles = create_virtual_project( + self.project, + config.profile_name.text(), + config.default_filename.text(), + config.schema_keys(), + ).profiles + if imported is not None: + imported = { + page.name: imported.get(page.name, {}) for page in profiles[0].pages + } + elif config.source_path is None: + profiles = self.project.profiles + for profile in profiles: + target = (self.project.output_root / profile.default_file).resolve() + previous = self.editors.get(target) + if previous is not None and (config.builder or config.source_path is not None): + raise ValueError(f"'{target.name}' is already open") + values = ( + previous.values() + if previous + else ( + imported + if imported is not None + else load_profile(self.project, profile, dimensions) + ) + ) + if previous is None: + values = overlay_values(values, self.initial_values.get(profile.key, {})) + editor = ProfileTab( + self.project, + profile, + values, + dimensions, + self, + fit_arrays=previous is not None and previous.dimensions != dimensions, + ) + editor.dimensions = dict(dimensions) + editor.saved_dimensions = dict( + previous.saved_dimensions if previous else dimensions + ) + editor.saved_values = copy.deepcopy( + previous.saved_values if previous else editor.values() + ) + prepared.append(editor) + for editor in prepared: + self._put_editor(editor) + self.dimensions = dimensions + if config.builder or config.source_path is not None: + self.config_tabs.discard(config) + self.tabs.removeTab(self.tabs.indexOf(config)) + config.deleteLater() + except (OSError, ValueError, KeyError) as exc: + for editor in prepared: + editor.deleteLater() + QMessageBox.critical(self, "Invalid configuration", str(exc)) + + def _put_editor(self, editor: ProfileTab) -> None: + self.tabs.blockSignals(True) + path = (self.project.output_root / editor.profile.default_file).resolve() + old = self.editors.get(path) + index = self.tabs.indexOf(old) if old else self.tabs.indexOf(self.plus_tab) + if old: + self.tabs.removeTab(index) + old.deleteLater() + self.editors[path] = editor + self.tabs.insertTab(index, editor, editor.profile.title) + self.tabs.setCurrentWidget(editor) + self.tabs.blockSignals(False) + editor.save.clicked.connect(lambda: self._save([editor])) + editor.cancel.clicked.connect(lambda: self._cancel_profile(editor)) + + def _cancel_profile(self, editor: ProfileTab) -> None: + replacement = ProfileTab( + self.project, editor.profile, editor.saved_values, editor.saved_dimensions, self + ) + replacement.dimensions = dict(editor.saved_dimensions) + replacement.saved_dimensions = dict(editor.saved_dimensions) + replacement.saved_values = copy.deepcopy(editor.saved_values) + self._put_editor(replacement) + + def _save(self, editors: list[ProfileTab]) -> None: + try: + values = [(editor, editor.values()) for editor in editors] + save_profiles( + self.project, + [(editor.profile, value, editor.dimensions) for editor, value in values], + ) + for editor, value in values: + editor.saved_values = copy.deepcopy(value) + editor.saved_dimensions = dict(editor.dimensions) + except (OSError, ValueError, KeyError) as exc: + QMessageBox.critical(self, "Save namelist", str(exc)) + + def _save_all(self) -> None: + self._save(list(self.editors.values())) + + def _restore_all(self) -> None: + for editor in self.editors.values(): + editor.restore_all() + + +def launch_gui( + schemas_dir: Path | str | None = None, + output_dir: Path | str | None = None, + file_profiles: Mapping[str, list[str]] | None = None, + initial_values: Mapping[str, Any] | None = None, + initial_dimensions: Mapping[str, int] | None = None, +) -> int: + """Launch independently or reuse the caller's QApplication.""" + project = load_project(schemas_dir, output_dir, file_profiles) + application = QApplication.instance() + owns_application = application is None + if application is None: + application = QApplication(sys.argv[:1]) + application.setApplicationName("nml-tools") + dialog = ConfigurationDialog( + project, initial_values=initial_values, initial_dimensions=initial_dimensions + ) + if not owns_application: + _exec(dialog) + return 0 + dialog.show() + method = getattr(application, "exec", None) or application.exec_ + return int(method()) diff --git a/src/nml_tools/gui/arrays.py b/src/nml_tools/gui/arrays.py new file mode 100644 index 0000000..798e6fa --- /dev/null +++ b/src/nml_tools/gui/arrays.py @@ -0,0 +1,257 @@ +"""Array shape, label, and display-order helpers for the GUI.""" + +from __future__ import annotations + +import copy +import math +from collections.abc import Iterator, Mapping, Sequence +from typing import Any, cast + + +def resolve_shape( + schema: Mapping[str, Any], + sizes: Mapping[str, int], + existing: Any = None, +) -> tuple[int, ...]: + """Resolve ``x-fortran-shape`` using configured constants and dimensions.""" + raw = schema.get("x-fortran-shape") + dimensions = raw if isinstance(raw, list) else [raw] + existing_shape = _list_shape(existing) + result: list[int] = [] + for axis, dimension in enumerate(dimensions): + if isinstance(dimension, bool): + raise ValueError("array shape must not contain booleans") + if isinstance(dimension, int): + value = dimension + elif isinstance(dimension, str): + token = dimension.strip() + if token == ":": + value = existing_shape[axis] if axis < len(existing_shape) else 1 + else: + try: + value = int(token) + except ValueError as exc: + value = sizes.get(token.lower(), 0) + if not value: + raise ValueError(f"unknown array dimension '{dimension}'") from exc + else: + raise ValueError("array property must define 'x-fortran-shape'") + if value <= 0: + raise ValueError("array dimensions must be positive") + result.append(value) + if not result: + raise ValueError("array property must define 'x-fortran-shape'") + return tuple(result) + + +def flex_tail_dims(schema: Mapping[str, Any], rank: int) -> int: + """Return the validated number of flexible trailing dimensions.""" + raw = schema.get("x-fortran-flex-tail-dims", 0) + if isinstance(raw, bool) or not isinstance(raw, int): + raise ValueError("'x-fortran-flex-tail-dims' must be an integer") + if not 0 <= raw <= rank: + raise ValueError("'x-fortran-flex-tail-dims' must be between zero and the array rank") + return raw + + +def array_shape(value: Any) -> tuple[int, ...]: + """Return the non-empty rectangular shape of a array value.""" + shape = _list_shape(value) + if not shape: + raise ValueError("array values must be non-empty and rectangular") + return shape + + +def validate_array_shape( + schema: Mapping[str, Any], sizes: Mapping[str, int], value: Any +) -> tuple[int, ...]: + """Validate a saved array shape and return its actual shape.""" + actual = array_shape(value) + declared = resolve_shape(schema, sizes, value) + if len(actual) != len(declared): + raise ValueError(f"array rank {len(actual)} does not match declared rank {len(declared)}") + flexible = flex_tail_dims(schema, len(declared)) + fixed = len(declared) - flexible + if actual[:fixed] != declared[:fixed]: + raise ValueError(f"array shape {actual} does not match declared shape {declared}") + if flexible and any(actual[index] > declared[index] for index in range(fixed, len(declared))): + raise ValueError(f"array shape {actual} exceeds declared shape {declared}") + if not flexible and actual != declared: + raise ValueError(f"array shape {actual} does not match declared shape {declared}") + return actual + + +def axis_labels(schema: Mapping[str, Any], axis: int, extent: int) -> list[str] | None: + """Return labels for a one-based array *axis*, validating their count.""" + metadata = schema.get("x-nml-tools-ui", {}) + if metadata is None: + return None + if not isinstance(metadata, Mapping): + raise ValueError("'x-nml-tools-ui' must be an object") + axes = metadata.get("axes", {}) + if not isinstance(axes, Mapping): + raise ValueError("'x-nml-tools-ui.axes' must be an object") + raw = axes.get(str(axis), axes.get(axis)) + if raw is None: + return None + if not isinstance(raw, Mapping): + raise ValueError(f"array UI axis {axis} must be an object") + title = raw.get("title") + if title is not None and not isinstance(title, str): + raise ValueError(f"array UI axis {axis} title must be a string") + labels = raw.get("labels") + template = raw.get("label-template") + if labels is not None and template is not None: + raise ValueError(f"array UI axis {axis} cannot define labels and label-template") + if labels is not None: + if not isinstance(labels, list) or not all(isinstance(item, str) for item in labels): + raise ValueError(f"array UI axis {axis} labels must be a list of strings") + if len(labels) != extent: + raise ValueError( + f"array UI axis {axis} defines {len(labels)} labels for extent {extent}" + ) + return list(labels) + if template is not None: + if not isinstance(template, str) or not template: + raise ValueError(f"array UI axis {axis} label-template must be a string") + try: + return [template.format(index=index) for index in range(1, extent + 1)] + except (KeyError, ValueError) as exc: + raise ValueError(f"array UI axis {axis} label-template must use '{{index}}'") from exc + return None + + +def table_axes(schema: Mapping[str, Any], rank: int) -> tuple[int, int] | None: + """Return zero-based ``(row, column)`` axes for a two-dimensional display.""" + metadata = schema.get("x-nml-tools-ui", {}) + if not isinstance(metadata, Mapping): + raise ValueError("'x-nml-tools-ui' must be an object") + table = metadata.get("table") + if table is None: + return (0, 1) if rank == 2 else None + if not isinstance(table, Mapping): + raise ValueError("'x-nml-tools-ui.table' must be an object") + row = table.get("row-axis") + column = table.get("column-axis") + if isinstance(row, bool) or not isinstance(row, int): + raise ValueError("array UI table row-axis must be an integer") + if isinstance(column, bool) or not isinstance(column, int): + raise ValueError("array UI table column-axis must be an integer") + if row == column or not 1 <= row <= rank or not 1 <= column <= rank: + raise ValueError("array UI table axes must be distinct valid one-based axes") + if rank != 2: + raise ValueError("array UI table orientation currently requires a rank-two array") + return row - 1, column - 1 + + +def display_array(value: Any, schema: Mapping[str, Any]) -> Any: + """Return a NumPy array ordered for the configured table display.""" + import numpy as np + + data = np.asarray(value) + if data.ndim == 1: + return data.reshape((1, data.shape[0])) + axes = table_axes(schema, data.ndim) + if axes is not None and axes != (0, 1): + return np.transpose(data, axes) + return data + + +def canonical_array(value: Any, schema: Mapping[str, Any], rank: int) -> list[Any]: + """Convert a displayed NumPy array back to canonical Fortran-axis order.""" + import numpy as np + + data = np.asarray(value) + if rank == 1: + data = data.reshape((-1,)) + else: + axes = table_axes(schema, rank) + if axes is not None and axes != (0, 1): + data = np.transpose(data, np.argsort(axes)) + return cast(list[Any], _native_value(data.tolist())) + + +def initial_array( + schema: Mapping[str, Any], + sizes: Mapping[str, int], + value: Any, + leaf_default: Any, + *, + strict: bool = False, +) -> list[Any]: + """Fit a saved/default/example value to the resolved canonical shape.""" + shape = resolve_shape(schema, sizes, value) + result = cast(list[Any], _filled(shape, leaf_default)) + if value is None: + return result + if not isinstance(value, list): + if strict: + raise ValueError("saved array values must be arrays") + return cast(list[Any], _filled(shape, value)) + current_shape = _list_shape(value) + if strict: + validate_array_shape(schema, sizes, value) + return copy.deepcopy(value) + if current_shape and flex_tail_dims(schema, len(shape)): + try: + validate_array_shape(schema, sizes, value) + except ValueError: + pass + else: + return copy.deepcopy(value) + if current_shape == shape: + return copy.deepcopy(value) + + flat = list(_flatten(value)) + if not flat: + return result + if len(shape) > 1 and len(flat) == shape[0]: + return _broadcast_first_axis(flat, shape) + if len(flat) == math.prod(shape): + return _reshape(flat, shape) + return _reshape([flat[index % len(flat)] for index in range(math.prod(shape))], shape) + + +def _broadcast_first_axis(values: list[Any], shape: tuple[int, ...]) -> list[Any]: + tail = shape[1:] + return [_filled(tail, value) for value in values] + + +def _filled(shape: Sequence[int], value: Any) -> Any: + if not shape: + return copy.deepcopy(value) + return [_filled(shape[1:], value) for _ in range(shape[0])] + + +def _reshape(values: Sequence[Any], shape: tuple[int, ...]) -> list[Any]: + iterator = iter(values) + + def build(remaining: tuple[int, ...]) -> Any: + if len(remaining) == 1: + return [copy.deepcopy(next(iterator)) for _ in range(remaining[0])] + return [build(remaining[1:]) for _ in range(remaining[0])] + + return cast(list[Any], build(shape)) + + +def _flatten(value: Any) -> Iterator[Any]: + if isinstance(value, list): + for child in value: + yield from _flatten(child) + else: + yield value + + +def _list_shape(value: Any) -> tuple[int, ...]: + if not isinstance(value, list) or not value: + return () + first = _list_shape(value[0]) + if any(_list_shape(item) != first for item in value[1:]): + return () + return (len(value), *first) + + +def _native_value(value: Any) -> Any: + if isinstance(value, list): + return [_native_value(item) for item in value] + return value.item() if hasattr(value, "item") else value diff --git a/src/nml_tools/gui/fields.py b/src/nml_tools/gui/fields.py new file mode 100644 index 0000000..57bf14a --- /dev/null +++ b/src/nml_tools/gui/fields.py @@ -0,0 +1,478 @@ +"""Schema-driven Qt field widgets.""" + +from __future__ import annotations + +import copy +import math +from collections.abc import Mapping +from typing import Any, cast + +from qtpy.QtWidgets import ( + QCheckBox, + QComboBox, + QFormLayout, + QGroupBox, + QHBoxLayout, + QLabel, + QLineEdit, + QMessageBox, + QPushButton, + QWidget, +) + +from .arrays import ( + array_shape, + axis_labels, + canonical_array, + display_array, + flex_tail_dims, + initial_array, + resolve_shape, + table_axes, +) +from .model import MISSING, overlay_values, suggestion + + +def _exec(dialog: Any) -> int: + method = getattr(dialog, "exec", None) + if method is None: + method = dialog.exec_ + return int(method()) + + +def _accepted(dialog: Any) -> int: + value = getattr(dialog, "Accepted", None) + value = value if value is not None else dialog.DialogCode.Accepted + return int(getattr(value, "value", value)) + + +def _derived_array_editor(editor_type: type[Any], parent: QWidget) -> Any: + # guidata's fixed-size record handler cannot commit (field, *indices) keys. + class DerivedArrayEditor(editor_type): # type: ignore[misc] + def accept(self) -> None: + for (name, *indices), value in self._data.current_changes.items(): + self._data.get_array()[name][tuple(indices)] = value + self._data.current_changes.clear() + super().accept() + + return DerivedArrayEditor(parent) + + +class ScalarField(QWidget): + def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None = None): + super().__init__(parent) + self.schema = schema + layout = QHBoxLayout(self) + layout.setContentsMargins(0, 0, 0, 0) + enum = schema.get("enum") + kind = schema.get("type") + control: QComboBox | QCheckBox | QLineEdit + if isinstance(enum, list) and enum: + combo = QComboBox(self) + for item in enum: + combo.addItem(str(item), item) + control = combo + elif kind == "boolean": + control = QCheckBox(self) + else: + control = QLineEdit(self) + self.control = control + layout.addWidget(control) + self.set_value(value) + + def set_value(self, value: Any) -> None: + if isinstance(self.control, QComboBox): + index = self.control.findData(value) + self.control.setCurrentIndex(max(index, 0)) + elif isinstance(self.control, QCheckBox): + self.control.setChecked(bool(value)) + else: + self.control.setText(str(value)) + + def value(self) -> Any: + if isinstance(self.control, QComboBox): + return self.control.currentData() + if isinstance(self.control, QCheckBox): + return self.control.isChecked() + text = self.control.text() + kind = self.schema.get("type") + try: + if kind == "integer": + return int(text) + if kind == "number": + value = float(text) + if not math.isfinite(value): + raise ValueError + return value + except ValueError as exc: + raise ValueError(f"'{text}' is not a valid {kind}") from exc + return text + + def reset(self, sizes: Mapping[str, int]) -> None: + self.set_value(suggestion(self.schema, sizes)) + + +class ObjectField(QGroupBox): + def __init__( + self, + schema: Mapping[str, Any], + value: Any, + sizes: Mapping[str, int], + parent: QWidget | None = None, + *, + fit_arrays: bool = False, + ): + super().__init__(str(schema.get("x-fortran-type", "Derived value")), parent) + self.schema = schema + self.sizes = sizes + properties = schema.get("properties") + if not isinstance(properties, Mapping): + raise ValueError("derived field must define object 'properties'") + source = suggestion(schema, sizes) + if isinstance(value, Mapping): + source = overlay_values(source, value) + required = {item.lower() for item in schema.get("required", []) if isinstance(item, str)} + layout = QFormLayout(self) + self.rows: dict[str, FieldRow] = {} + for name, child in properties.items(): + if not isinstance(name, str) or not isinstance(child, Mapping): + continue + child_value = source.get(name, MISSING) if isinstance(source, Mapping) else MISSING + is_required = name.lower() in required + row = FieldRow(name, child, child_value, sizes, self, fit_arrays=fit_arrays) + layout.addRow(_field_label(name, child, is_required), row) + self.rows[name] = row + + def value(self) -> dict[str, Any]: + result: dict[str, Any] = {} + for name, row in self.rows.items(): + value = row.value() + if value is not MISSING: + result[name] = value + return result + + def reset(self, sizes: Mapping[str, int]) -> None: + defaults = suggestion(self.schema, sizes) + for name, row in self.rows.items(): + row.set_value(defaults[name], sizes) + + +class ArrayField(QWidget): + def __init__( + self, + name: str, + schema: Mapping[str, Any], + value: Any, + sizes: Mapping[str, int], + parent: QWidget | None = None, + *, + fit_existing: bool = False, + ): + super().__init__(parent) + self.name = name + self.schema = schema + self.sizes = sizes + saved = value is not MISSING + candidate = suggestion(schema, sizes) if not saved else value + items = schema.get("items") + if not isinstance(items, Mapping): + raise ValueError("array field must define object 'items'") + self.items = items + self._value = initial_array( + schema, + sizes, + candidate, + suggestion(items, sizes), + strict=saved and not fit_existing, + ) + layout = QHBoxLayout(self) + layout.setContentsMargins(0, 0, 0, 0) + self.summary = QLabel(self) + self.inline: Any = None + self.button = QPushButton("Edit array…", self) + self.button.clicked.connect(self._edit) + layout.addWidget(self.summary, 1) + layout.addWidget(self.button) + self._update_summary() + + def value(self) -> list[Any]: + if self.inline is not None: + return [self.inline.value()] + return copy.deepcopy(self._value) + + def reset(self, sizes: Mapping[str, int]) -> None: + self.sizes = sizes + self._value = suggestion(self.schema, sizes) + self._update_summary() + + def _update_summary(self) -> None: + shape = array_shape(self._value) + if self.inline is not None: + self.layout().removeWidget(self.inline) + self.inline.deleteLater() + self.inline = None + if shape == (1,): + self.inline = _field_widget(self.name, self.items, self._value[0], self.sizes, self) + self.layout().insertWidget(0, self.inline, 1) + raw = self.schema.get("x-fortran-shape") + resizable = raw == ":" or (isinstance(raw, list) and ":" in raw) + resizable = resizable or flex_tail_dims(self.schema, len(shape)) > 0 + self.summary.setVisible(self.inline is None) + self.button.setVisible(self.inline is None or resizable) + self.button.setText("Resize…" if self.inline is not None else "Edit array…") + self.summary.setText("×".join(str(value) for value in shape)) + + def _edit(self) -> None: + try: + import numpy as np + from guidata.widgets.arrayeditor import ArrayEditor # type: ignore[import-untyped] + + self._value = self.value() + rank = len(resolve_shape(self.schema, self.sizes, self._value)) + derived = self.items.get("type") == "object" + canonical = self._structured_array(np) if derived else self._intrinsic_array(np) + displayed = display_array(canonical, self.schema) + xlabels, ylabels = self._display_labels(displayed.shape, rank) + editor = _derived_array_editor(ArrayEditor, self) if derived else ArrayEditor(self) + raw_shape = self.schema.get("x-fortran-shape") + deferred = raw_shape == ":" or (isinstance(raw_shape, list) and ":" in raw_shape) + if not editor.setup_and_check( + displayed, + str(self.schema.get("title", self.name)), + xlabels=xlabels, + ylabels=ylabels, + variable_size=flex_tail_dims(self.schema, rank) > 0 or deferred, + ): + return + if _exec(editor) != _accepted(editor): + return + edited = editor.get_value() + if derived: + self._value = self._objects_from_structured(edited, rank, np) + else: + self._value = canonical_array(edited, self.schema, rank) + self._update_summary() + except (ImportError, RuntimeError, TypeError, ValueError) as exc: + QMessageBox.critical(self, "Array editor", str(exc)) + + def _intrinsic_array(self, np: Any) -> Any: + kind = self.items.get("type") + if not isinstance(kind, str): + raise ValueError("array items must define a string type") + dtype = { + "integer": np.int64, + "number": np.float64, + "boolean": np.bool_, + "string": "U1024", + }.get(kind) + if dtype is None: + raise ValueError(f"unsupported array item type '{kind}'") + return np.asarray(self._value, dtype=dtype) + + def _structured_array(self, np: Any) -> Any: + properties = self.items.get("properties") + if not isinstance(properties, Mapping): + raise ValueError("derived array items must define properties") + fields = [] + for name, child in properties.items(): + if not isinstance(name, str) or not isinstance(child, Mapping): + continue + child_kind = child.get("type") + if not isinstance(child_kind, str): + raise ValueError("derived components must define a string type") + dtype = { + "integer": np.int64, + "number": np.float64, + "boolean": np.bool_, + "string": "U1024", + }.get(child_kind) + if dtype is None: + raise ValueError(f"unsupported derived component type '{child_kind}'") + title = str(child.get("title", name)) + fields.append((name, dtype) if title == name else ((title, name), dtype)) + shape = array_shape(self._value) + result = np.empty(shape, dtype=np.dtype(fields)) + defaults = suggestion(self.items, self.sizes) + for index in np.ndindex(shape): + item = _nested_get(self._value, index) + if not isinstance(item, Mapping): + item = defaults + for name in result.dtype.names or (): + result[index][name] = item.get(name, defaults[name]) + return result + + def _objects_from_structured(self, value: Any, rank: int, np: Any) -> list[Any]: + data = np.asarray(value) + if rank == 1: + data = data.reshape((-1,)) + else: + axes = table_axes(self.schema, rank) + if axes is not None and axes != (0, 1): + data = np.transpose(data, np.argsort(axes)) + edited = _structured_to_objects(data) + defaults = suggestion(self.items, self.sizes) + return cast(list[Any], _preserve_omissions(edited, self._value, defaults)) + + def _display_labels( + self, displayed_shape: tuple[int, ...], rank: int + ) -> tuple[list[str] | None, list[str] | None]: + if rank == 1: + return axis_labels(self.schema, 1, displayed_shape[1]), None + if rank != 2: + return None, None + axes = table_axes(self.schema, rank) or (0, 1) + return ( + axis_labels(self.schema, axes[1] + 1, displayed_shape[1]), + axis_labels(self.schema, axes[0] + 1, displayed_shape[0]), + ) + + +class FieldRow(QWidget): + def __init__( + self, + name: str, + schema: Mapping[str, Any], + value: Any, + sizes: Mapping[str, int], + parent: QWidget | None = None, + *, + fit_arrays: bool = False, + ): + super().__init__(parent) + self.name = name + self.schema = schema + self.sizes = sizes + layout = QHBoxLayout(self) + layout.setContentsMargins(0, 0, 0, 0) + initial = value + if value is MISSING and schema.get("type") not in {"array", "object"}: + initial = suggestion(schema, sizes) + self.field = _field_widget(name, schema, initial, sizes, self, fit_arrays=fit_arrays) + description = schema.get("description") + if isinstance(description, str): + self.field.setToolTip(description.strip()) + layout.addWidget(self.field, 1) + + def value(self) -> Any: + return self.field.value() + + def reset(self, sizes: Mapping[str, int]) -> None: + self.set_value(suggestion(self.schema, sizes), sizes) + + def set_value(self, value: Any, sizes: Mapping[str, int]) -> None: + self.sizes = sizes + replacement = _field_widget(self.name, self.schema, value, sizes, self) + replacement.setToolTip(self.field.toolTip()) + self.layout().replaceWidget(self.field, replacement) + self.field.deleteLater() + self.field = replacement + + +class NamelistForm(QWidget): + """Editable form for one namelist schema.""" + + def __init__( + self, + schema: Mapping[str, Any], + values: Mapping[str, Any] | None, + sizes: Mapping[str, int], + parent: QWidget | None = None, + *, + fit_arrays: bool = False, + ): + super().__init__(parent) + self.schema = schema + self.sizes = sizes + properties = schema.get("properties") + if not isinstance(properties, Mapping): + raise ValueError("namelist schema must define object 'properties'") + source = values or {} + required = {item.lower() for item in schema.get("required", []) if isinstance(item, str)} + layout = QFormLayout(self) + self.rows: dict[str, FieldRow] = {} + for name, child in properties.items(): + if not isinstance(name, str) or not isinstance(child, Mapping): + continue + is_required = name.lower() in required + row = FieldRow( + name, + child, + source.get(name, MISSING), + sizes, + self, + fit_arrays=fit_arrays, + ) + layout.addRow(_field_label(name, child, is_required), row) + self.rows[name] = row + + def values(self) -> dict[str, Any]: + result: dict[str, Any] = {} + for name, row in self.rows.items(): + value = row.value() + if value is not MISSING: + result[name] = value + return result + + def reset(self) -> None: + for row in self.rows.values(): + row.reset(self.sizes) + + +def _field_widget( + name: str, + schema: Mapping[str, Any], + value: Any, + sizes: Mapping[str, int], + parent: QWidget, + *, + fit_arrays: bool = False, +) -> Any: + kind = schema.get("type") + if kind == "array": + return ArrayField(name, schema, value, sizes, parent, fit_existing=fit_arrays) + if kind == "object": + return ObjectField(schema, value, sizes, parent, fit_arrays=fit_arrays) + return ScalarField(schema, value, parent) + + +def _field_label(name: str, schema: Mapping[str, Any], required: bool) -> str: + title = schema.get("title") + label = f"{title} ({name})" if isinstance(title, str) and title.strip() else name + return f"{label} *" if required else label + + +def _nested_get(value: Any, indices: tuple[int, ...]) -> Any: + for index in indices: + value = value[index] + return value + + +def _structured_to_objects(data: Any) -> list[Any]: + names = data.dtype.names or () + + def build(axis: int, prefix: tuple[int, ...]) -> Any: + if axis == data.ndim: + record = data[prefix] + return {name: _numpy_scalar(record[name]) for name in names} + return [build(axis + 1, (*prefix, index)) for index in range(data.shape[axis])] + + return cast(list[Any], build(0, ())) + + +def _numpy_scalar(value: Any) -> Any: + return value.item() if hasattr(value, "item") else value + + +def _preserve_omissions(edited: Any, original: Any, defaults: Any) -> Any: + if isinstance(edited, list) and isinstance(original, list): + return [ + _preserve_omissions(item, original[index], defaults) if index < len(original) else item + for index, item in enumerate(edited) + ] + if isinstance(edited, Mapping) and isinstance(original, Mapping): + return { + name: value + for name, value in edited.items() + if name in original or not isinstance(defaults, Mapping) or value != defaults.get(name) + } + return edited diff --git a/src/nml_tools/gui/model.py b/src/nml_tools/gui/model.py new file mode 100644 index 0000000..cb89863 --- /dev/null +++ b/src/nml_tools/gui/model.py @@ -0,0 +1,674 @@ +"""Qt-independent project loading and direct namelist persistence.""" + +from __future__ import annotations + +import copy +import math +import os +import tempfile +from collections.abc import Iterable +from dataclasses import dataclass, replace +from itertools import product +from pathlib import Path +from typing import Any, Mapping + +import click + +from .._namelist_eval import EvaluatedGroup, LeafState, evaluate_group +from .._namelist_parser import parse_namelist +from ..cli import ( + _iter_file_profiles, + _load_config_checked, + _load_constants, + _load_dimensions, + _load_namelist_registry, + _namelist_registry_by_key, +) +from ..codegen_fortran import _format_scalar_default +from ..schema import SchemaResolver +from ..validate import _scalar_constraints, _validate_scalar_value, validate_schema_defaults +from .arrays import flex_tail_dims, initial_array, resolve_shape, validate_array_shape + +MISSING = object() + + +def suggestion(schema: Mapping[str, Any], sizes: Mapping[str, int]) -> Any: + """Return the deterministic editable value used for an unset schema field.""" + examples = schema.get("examples") + if "default" in schema: + candidate = copy.deepcopy(schema["default"]) + elif isinstance(examples, list) and examples: + candidate = copy.deepcopy(examples[0]) + else: + candidate = MISSING + + kind = schema.get("type") + if kind == "array": + items = schema.get("items") + if not isinstance(items, Mapping): + raise ValueError("array field must define object 'items'") + leaf = suggestion(items, sizes) + if "default" in schema and isinstance(candidate, list): + shape = resolve_shape(schema, sizes) + count = math.prod(shape) + if schema.get("x-fortran-default-repeat"): + candidate = [candidate[index % len(candidate)] for index in range(count)] + elif "x-fortran-default-pad" in schema: + pad = schema["x-fortran-default-pad"] + pad = pad if isinstance(pad, list) else [pad] + candidate += [pad[index % len(pad)] for index in range(count - len(candidate))] + if len(candidate) != count: + raise ValueError("array default does not match its configured dimensions") + result = _filled(shape, leaf) + for index, coordinates in enumerate(product(*(range(size) for size in shape))): + if schema.get("x-fortran-default-order", "F").upper() == "F": + index = sum(c * math.prod(shape[:axis]) for axis, c in enumerate(coordinates)) + target = result + for coordinate in coordinates[:-1]: + target = target[coordinate] + target[coordinates[-1]] = copy.deepcopy(candidate[index]) + return result + if "default" in items: + candidate = MISSING + return initial_array(schema, sizes, None if candidate is MISSING else candidate, leaf) + if kind == "object": + raw = candidate if isinstance(candidate, Mapping) else {} + properties = schema.get("properties") + if not isinstance(properties, Mapping): + raise ValueError("derived field must define object 'properties'") + return { + name: copy.deepcopy(raw[name]) if name in raw else suggestion(child, sizes) + for name, child in properties.items() + if isinstance(name, str) and isinstance(child, Mapping) + } + if candidate is not MISSING: + return candidate + enum = schema.get("enum") + if isinstance(enum, list) and enum: + return copy.deepcopy(enum[0]) + if kind == "boolean": + return False + if kind == "integer": + minimum = schema.get("minimum") + return int(minimum) if isinstance(minimum, int) and not isinstance(minimum, bool) else 0 + if kind == "number": + minimum = schema.get("minimum") + return float(minimum) if isinstance(minimum, (int, float)) else 0.0 + if kind == "string": + return "" + raise ValueError(f"unsupported schema type '{kind}'") + + +@dataclass(frozen=True) +class NamelistPage: + """A configured namelist and its resolved schema.""" + + name: str + key: str + schema: dict[str, Any] + + +@dataclass(frozen=True) +class GuiProfile: + """An ordered file profile presented by the GUI.""" + + name: str + key: str + title: str + description: str | None + default_file: str + pages: tuple[NamelistPage, ...] + + +@dataclass(frozen=True) +class GuiProject: + """Resolved nml-tools project data needed by the GUI.""" + + root: Path + constants: dict[str, int] + default_dimensions: dict[str, int] + profiles: tuple[GuiProfile, ...] + output_dir: Path | None = None + namelists: tuple[NamelistPage, ...] = () + + @property + def output_root(self) -> Path: + """Return the directory used for namelist output.""" + return self.output_dir or self.root + + def profile(self, key: str) -> GuiProfile: + for profile in self.profiles: + if profile.key == key.lower(): + return profile + raise KeyError(key) + + +def load_project( + schemas_dir: Path | str | None = None, + output_dir: Path | str | None = None, + file_profiles: Mapping[str, list[str]] | None = None, +) -> GuiProject: + """Load schemas and profiles, using a separate output directory if given.""" + root = Path.cwd() if schemas_dir is None else Path(schemas_dir) + root = root.resolve() + output_root = root if output_dir is None else Path(output_dir).resolve() + config_path = root / "nml-config.toml" + if not config_path.is_file(): + raise RuntimeError(f"nml-config.toml was not found in {root}") + + try: + config, resolved_path = _load_config_checked(config_path) + constants, _ = _load_constants(config) + dimensions, _ = _load_dimensions(config, constants) + loaded = _load_namelist_registry(config, resolved_path.parent, SchemaResolver()) + registry = _namelist_registry_by_key(loaded) + configured_profiles = _iter_file_profiles(config, registry) + except click.ClickException as exc: + raise RuntimeError(exc.format_message()) from exc + except (OSError, ValueError) as exc: + raise RuntimeError(str(exc)) from exc + + namelists = tuple(NamelistPage(item.name, item.key, item.schema) for item in loaded) + pages_by_key = {page.key: page for page in namelists} + profiles: list[GuiProfile] = [] + output_paths: dict[Path, str] = {} + for configured in configured_profiles.values(): + target = (output_root / configured.default_file).resolve() + try: + target.relative_to(output_root) + except ValueError as exc: + raise RuntimeError( + f"file profile '{configured.name}' writes outside the output directory" + ) from exc + previous = output_paths.get(target) + if previous is not None: + raise RuntimeError( + f"file profiles '{previous}' and '{configured.name}' both write {target}" + ) + output_paths[target] = configured.name + pages = tuple(pages_by_key[key] for key in configured.namelists) + profiles.append( + GuiProfile( + name=configured.name, + key=configured.key, + title=configured.title or configured.name, + description=configured.description, + default_file=configured.default_file, + pages=pages, + ) + ) + + project = GuiProject( + root, + constants, + dimensions, + tuple(profiles), + output_root, + namelists, + ) + if file_profiles is None: + return project + if not isinstance(file_profiles, Mapping): + raise ValueError("file_profiles must map profile names to lists of namelist names") + selected = {} + for name, names in file_profiles.items(): + if not isinstance(name, str) or not isinstance(names, list): + raise ValueError("file_profiles must map profile names to lists of namelist names") + try: + profile = project.profile(name) + except KeyError as exc: + raise ValueError(f"unknown file profile '{name}'") from exc + if profile.key in selected: + raise ValueError(f"duplicate file profile '{name}'") + if not all(isinstance(item, str) for item in names): + raise ValueError(f"namelist names for '{name}' must be strings") + keys = {item.lower() for item in names} + unknown = keys - {page.key for page in profile.pages} + if unknown: + raise ValueError(f"profile '{name}' has no namelists: {', '.join(sorted(unknown))}") + selected[profile.key] = replace( + profile, pages=tuple(page for page in profile.pages if not keys or page.key in keys) + ) + if selected: + project = replace( + project, profiles=tuple(selected[p.key] for p in project.profiles if p.key in selected) + ) + return project + + +def create_virtual_project( + project: GuiProject, + name: str, + default_file: str, + namelist_keys: Iterable[str], +) -> GuiProject: + """Return one user-defined file profile using registered namelists.""" + if not isinstance(name, str): + raise ValueError("file profile name must be a string") + clean_name = name.strip() + if not clean_name: + raise ValueError("file profile name must not be empty") + if not isinstance(default_file, str): + raise ValueError("default file name must be a string") + clean_default = default_file.strip() + if not clean_default: + raise ValueError("default file name must not be empty") + + target = (project.output_root / clean_default).resolve() + try: + relative = target.relative_to(project.output_root) + except ValueError as exc: + raise ValueError("default file must be inside the output directory") from exc + if target == project.output_root: + raise ValueError("default file must name a file") + + available = {page.key: page for page in project.namelists} + selected: list[NamelistPage] = [] + seen: set[str] = set() + for raw_key in namelist_keys: + if not isinstance(raw_key, str) or not raw_key.strip(): + raise ValueError("selected namelist names must be non-empty strings") + key = raw_key.lower() + if key in seen: + raise ValueError(f"selected namelist '{raw_key}' is duplicated") + seen.add(key) + page = available.get(key) + if page is None: + raise ValueError(f"selected namelist '{raw_key}' is unknown") + selected.append(page) + if not selected: + raise ValueError("at least one namelist schema must be selected") + + profile = GuiProfile( + name=clean_name, + key=clean_name.lower(), + title=clean_name, + description=None, + default_file=str(relative), + pages=tuple(selected), + ) + return replace(project, profiles=(profile,)) + + +def _evaluated_group_values( + evaluated: EvaluatedGroup, + schema: Mapping[str, Any], + sizes: Mapping[str, int], +) -> dict[str, Any]: + properties = schema.get("properties", {}) + if not isinstance(properties, Mapping): + raise ValueError(f"schema for namelist '{evaluated.name}' has invalid properties") + result: dict[str, Any] = {} + for name, prop in properties.items(): + if not isinstance(name, str) or not isinstance(prop, Mapping): + continue + states = [ + (coordinates, component, state) + for (root, coordinates, component), state in evaluated.states.items() + if root == name.lower() and state.explicitly_assigned + ] + if not states: + continue + if prop.get("type") == "array": + result[name] = _evaluated_array(prop, states, sizes) + elif prop.get("type") == "object": + components = _component_names(prop) + result[name] = { + components[component]: _imported_scalar(state.value) + for _, component, state in states + if component in components + } + else: + result[name] = _imported_scalar(states[-1][2].value) + return result + + +def _evaluated_array( + schema: Mapping[str, Any], + states: list[tuple[tuple[int, ...], str | None, LeafState]], + sizes: Mapping[str, int], +) -> list[Any]: + items = schema.get("items") + if not isinstance(items, Mapping): + raise ValueError("array field must define object 'items'") + shape = list(resolve_shape(schema, sizes)) + flexible = flex_tail_dims(schema, len(shape)) + raw = schema["x-fortran-shape"] + raw = raw if isinstance(raw, list) else [raw] + for axis in range(len(shape)): + if raw[axis] != ":" and axis < len(shape) - flexible: + continue + used = [coordinates[axis] for coordinates, _, _ in states if coordinates] + if used: + shape[axis] = max(used) + result = suggestion({**schema, "x-fortran-shape": shape}, sizes) + components = _component_names(items) if items.get("type") == "object" else {} + for coordinates, component, state in states: + target = result + for coordinate in coordinates[:-1]: + target = target[coordinate - 1] + index = coordinates[-1] - 1 + value = _imported_scalar(state.value) + if component is None: + target[index] = value + elif component in components: + target[index][components[component]] = value + return result + + +def _component_names(schema: Mapping[str, Any]) -> dict[str, str]: + properties = schema.get("properties", {}) + if not isinstance(properties, Mapping): + return {} + return {name.lower(): name for name in properties if isinstance(name, str)} + + +def _filled(shape: tuple[int, ...], value: Any) -> Any: + if not shape: + return copy.deepcopy(value) + return [_filled(shape[1:], value) for _ in range(shape[0])] + + +def _imported_scalar(value: Any) -> Any: + return value.rstrip() if isinstance(value, str) else copy.deepcopy(value) + + +def _normalize_profile_values( + raw: Any, + profile: GuiProfile, + sizes: Mapping[str, int], +) -> dict[str, dict[str, Any]]: + if not isinstance(raw, Mapping): + raise ValueError(f"file profile '{profile.name}' values must be an object") + pages = {page.key: page for page in profile.pages} + values: dict[str, dict[str, Any]] = {} + seen_pages: set[str] = set() + for raw_name, fields in raw.items(): + if not isinstance(raw_name, str) or not isinstance(fields, Mapping): + raise ValueError(f"profile '{profile.name}' namelists must be named objects") + page = pages.get(raw_name.lower()) + if page is None: + raise ValueError(f"profile '{profile.name}' contains unknown namelist '{raw_name}'") + if page.key in seen_pages: + raise ValueError( + f"profile '{profile.name}' repeats namelist '{raw_name}' case-insensitively" + ) + seen_pages.add(page.key) + properties = page.schema.get("properties", {}) + if not isinstance(properties, Mapping): + raise ValueError(f"schema for namelist '{page.name}' has invalid properties") + canonical = { + str(name).lower(): (str(name), schema) + for name, schema in properties.items() + if isinstance(schema, Mapping) + } + normalized_fields: dict[str, Any] = {} + seen_fields: set[str] = set() + for raw_field, value in fields.items(): + if not isinstance(raw_field, str): + raise ValueError(f"namelist '{page.name}' field names must be strings") + field = canonical.get(raw_field.lower()) + if field is None: + raise ValueError(f"namelist '{page.name}' contains unknown field '{raw_field}'") + field_name, field_schema = field + field_key = field_name.lower() + if field_key in seen_fields: + raise ValueError( + f"namelist '{page.name}' repeats field '{raw_field}' case-insensitively" + ) + seen_fields.add(field_key) + normalized_fields[field_name] = _normalize_value( + value, + field_schema, + sizes, + f"{page.name}.{field_name}", + ) + values[page.name] = normalized_fields + return values + + +def _normalize_value( + value: Any, + schema: Mapping[str, Any], + sizes: Mapping[str, int], + path: str, +) -> Any: + kind = schema.get("type") + if kind == "array": + if not isinstance(value, list): + raise ValueError(f"'{path}' must be an array") + validate_array_shape(schema, sizes, value) + items = schema.get("items") + if not isinstance(items, Mapping): + raise ValueError(f"array '{path}' must define object items") + + def normalize_items(node: Any, indices: tuple[int, ...] = ()) -> Any: + if isinstance(node, list): + return [ + normalize_items(item, (*indices, index)) + for index, item in enumerate(node, start=1) + ] + suffix = "".join(f"[{index}]" for index in indices) + return _normalize_value(node, items, sizes, f"{path}{suffix}") + + return normalize_items(value) + if kind == "object": + if not isinstance(value, Mapping): + raise ValueError(f"'{path}' must be an object") + properties = schema.get("properties") + if not isinstance(properties, Mapping): + raise ValueError(f"derived value '{path}' must define properties") + canonical = { + str(name).lower(): (str(name), child) + for name, child in properties.items() + if isinstance(child, Mapping) + } + result: dict[str, Any] = {} + seen: set[str] = set() + for raw_name, child_value in value.items(): + if not isinstance(raw_name, str): + raise ValueError(f"derived value '{path}' component names must be strings") + child = canonical.get(raw_name.lower()) + if child is None: + raise ValueError(f"derived value '{path}' contains unknown component '{raw_name}'") + child_name, child_schema = child + child_key = child_name.lower() + if child_key in seen: + raise ValueError( + f"derived value '{path}' repeats component '{raw_name}' case-insensitively" + ) + seen.add(child_key) + result[child_name] = _normalize_value( + child_value, + child_schema, + sizes, + f"{path}.{child_name}", + ) + return result + constraints = _scalar_constraints(path, schema, str(kind), dict(sizes), None) + _validate_scalar_value(path, value, constraints) + return copy.deepcopy(value) + + +def _normalize_dimensions(dimensions: Mapping[str, int], project: GuiProject) -> dict[str, int]: + if not isinstance(dimensions, Mapping): + raise ValueError("dimensions must map names to positive integers") + result = dict(project.default_dimensions) + for name, value in dimensions.items(): + if not isinstance(name, str): + raise ValueError("dimension names must be strings") + key = name.lower() + if key not in result: + raise ValueError(f"unknown dimension '{name}'") + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise ValueError(f"dimension '{name}' must be a positive integer") + result[key] = value + return result + + +def overlay_values(base: Mapping[str, Any], override: Mapping[str, Any]) -> dict[str, Any]: + """Overlay supplied fields/components without discarding sibling values.""" + result = copy.deepcopy(dict(base)) + for name, value in override.items(): + if isinstance(value, Mapping) and isinstance(result.get(name), Mapping): + result[name] = overlay_values(result[name], value) + else: + result[name] = copy.deepcopy(value) + return result + + +def _profile_path(project: GuiProject, profile: GuiProfile) -> Path: + path = (project.output_root / profile.default_file).resolve() + if path == project.output_root or project.output_root not in path.parents: + raise ValueError("profile output must be a file inside the output directory") + return path + + +def _parsed_groups(path: Path) -> tuple[str, dict[str, Any]]: + text = path.read_text(encoding="utf-8") + parsed = parse_namelist(text, source=str(path)) + groups = {} + for group in parsed.groups: + key = group.name.lower() + if key in groups: + raise ValueError(f"namelist '{group.name}' appears multiple times in {path}") + groups[key] = group + if not groups: + raise ValueError(f"namelist file '{path}' contains no namelist groups") + return text, groups + + +def load_profile( + project: GuiProject, + profile: GuiProfile, + dimensions: Mapping[str, int], + path: Path | None = None, +) -> dict[str, Any]: + """Read selected groups; absent fields are supplied by the form's defaults.""" + dimensions = _normalize_dimensions(dimensions, project) + for page in profile.pages: + validate_schema_defaults(page.schema, constants=project.constants, dimensions=dimensions) + path = _profile_path(project, profile) if path is None else path + if not path.exists(): + return {} + _, groups = _parsed_groups(path) + sizes = {**project.constants, **dimensions} + result = {} + for page in profile.pages: + group = groups.get(page.key) + if group is not None: + evaluated = evaluate_group( + group, + page.schema, + source=str(path), + constants=project.constants, + dimensions=dimensions, + ) + result[page.name] = _evaluated_group_values(evaluated, page.schema, sizes) + return result + + +def import_profile( + project: GuiProject, path: Path, dimensions: Mapping[str, int] +) -> tuple[GuiProfile, dict[str, Any]]: + """Import a namelist file, checking every group against the project registry.""" + _, groups = _parsed_groups(path) + unknown = groups.keys() - {page.key for page in project.namelists} + if unknown: + raise ValueError( + f"Namelists {', '.join(sorted(unknown))} are not part of this nml-config.toml" + ) + matches = [ + profile + for profile in project.profiles + if Path(profile.default_file).name.casefold() == path.name.casefold() + ] + profile = ( + matches[0] + if len(matches) == 1 + else create_virtual_project(project, path.stem, path.name, groups).profiles[0] + ) + return profile, load_profile(project, profile, dimensions, path) + + +def _assignments(name: str, value: Any, schema: Mapping[str, Any]) -> Iterable[str]: + if schema["type"] == "array": + + def elements(node: Any, indices: tuple[int, ...] = ()) -> Iterable[str]: + if isinstance(node, list): + for index, child in enumerate(node, 1): + yield from elements(child, (*indices, index)) + else: + suffix = ",".join(map(str, indices)) + yield from _assignments(f"{name}({suffix})", node, schema["items"]) + + yield from elements(value) + elif schema["type"] == "object": + for component, child in value.items(): + yield from _assignments(f"{name}%{component}", child, schema["properties"][component]) + else: + category = "real" if schema["type"] == "number" else schema["type"] + yield f" {name} = {_format_scalar_default(value, None, category)}" + + +def render_profile( + project: GuiProject, + profile: GuiProfile, + values: Mapping[str, Any], + dimensions: Mapping[str, int], +) -> dict[str, str]: + """Render and validate individual groups, retaining explicit Fortran indices.""" + dimensions = _normalize_dimensions(dimensions, project) + normalized = _normalize_profile_values(values, profile, {**project.constants, **dimensions}) + rendered = {} + for page in profile.pages: + if page.name not in normalized: + continue + lines = [f"&{page.name}"] + for name, value in normalized[page.name].items(): + lines.extend(_assignments(name, value, page.schema["properties"][name])) + text = "\n".join([*lines, "/", ""]) + evaluate_group( + parse_namelist(text).groups[0], + page.schema, + constants=project.constants, + dimensions=dimensions, + ) + rendered[page.key] = text + return rendered + + +def save_profiles( + project: GuiProject, + updates: Iterable[tuple[GuiProfile, Mapping[str, Any], Mapping[str, int]]], +) -> None: + """Validate all updates, then replace files without removing unselected groups.""" + outputs: dict[Path, str] = {} + for profile, values, dimensions in updates: + path = _profile_path(project, profile) + if path in outputs: + raise ValueError(f"multiple open profiles write to '{path}'") + rendered = render_profile(project, profile, values, dimensions) + text, groups = _parsed_groups(path) if path.exists() else ("", {}) + for key, group in reversed(list(groups.items())): + if key in rendered: + replacement = rendered.pop(key).rstrip("\n") + text = text[: group.span.start.offset] + replacement + text[group.span.end.offset :] + for replacement in rendered.values(): + text += ("\n" if text and not text.endswith("\n") else "") + replacement + outputs[path] = text + for path, text in outputs.items(): + _atomic_write(path, text) + + +def _atomic_write(path: Path, content: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) + temporary = Path(temporary_name) + try: + with os.fdopen(descriptor, "w", encoding="utf-8", newline="\n") as handle: + handle.write(content) + os.replace(temporary, path) + except Exception: + temporary.unlink(missing_ok=True) + raise diff --git a/tests/test_cli_gui.py b/tests/test_cli_gui.py new file mode 100644 index 0000000..4efdced --- /dev/null +++ b/tests/test_cli_gui.py @@ -0,0 +1,36 @@ +"""GUI CLI wiring without starting a Qt event loop.""" + +import subprocess +import sys + +from click.testing import CliRunner + +import nml_tools.gui +from nml_tools.cli import cli + + +def test_gui_command_paths_and_help(tmp_path, monkeypatch): + calls = [] + monkeypatch.setattr(nml_tools.gui, "launch_gui", lambda *args: calls.append(args) or 0) + result = CliRunner().invoke(cli, ["gui", "-i", str(tmp_path), "-o", str(tmp_path / "out")]) + assert result.exit_code == 0 + assert calls == [(tmp_path, tmp_path / "out")] + help_text = CliRunner().invoke(cli, ["gui", "--help"]).output + assert "--input-path" in help_text and "--output-path" in help_text + assert "--fetch-values" not in help_text + + +def test_gui_import_is_lazy(): + subprocess.run( + [ + sys.executable, + "-c", + ( + "import sys, nml_tools.gui; " + "assert 'qtpy' not in sys.modules; " + "assert 'guidata' not in sys.modules; " + "assert 'numpy' not in sys.modules" + ), + ], + check=True, + ) diff --git a/tests/test_gui_arrays.py b/tests/test_gui_arrays.py new file mode 100644 index 0000000..eef955f --- /dev/null +++ b/tests/test_gui_arrays.py @@ -0,0 +1,141 @@ +"""Tests for Qt-independent GUI array metadata and data transforms.""" + +from __future__ import annotations + +import pytest + +from nml_tools.gui.arrays import ( + axis_labels, + canonical_array, + display_array, + initial_array, + resolve_shape, + table_axes, +) + + +def test_resolve_shape_supports_literals_names_and_saved_flexible_axes() -> None: + assert resolve_shape({"x-fortran-shape": [5, "N_DOMAINS"]}, {"n_domains": 3}) == (5, 3) + assert resolve_shape({"x-fortran-shape": ":"}, {}, [1, 2, 3]) == (3,) + + with pytest.raises(ValueError, match="unknown array dimension"): + resolve_shape({"x-fortran-shape": "missing"}, {}) + + +def test_axis_labels_support_explicit_values_and_index_templates() -> None: + schema = { + "x-nml-tools-ui": { + "axes": { + "1": {"labels": ["Lower", "Upper", "Value"]}, + "2": {"label-template": "Domain {index}"}, + } + } + } + + assert axis_labels(schema, 1, 3) == ["Lower", "Upper", "Value"] + assert axis_labels(schema, 2, 2) == ["Domain 1", "Domain 2"] + assert axis_labels(schema, 3, 4) is None + + with pytest.raises(ValueError, match="2 labels for extent 3"): + axis_labels( + {"x-nml-tools-ui": {"axes": {"1": {"labels": ["a", "b"]}}}}, + 1, + 3, + ) + + +def test_table_axes_uses_one_based_schema_metadata() -> None: + schema = { + "x-nml-tools-ui": { + "table": {"row-axis": 2, "column-axis": 1}, + } + } + + assert table_axes(schema, 2) == (1, 0) + assert table_axes({}, 2) == (0, 1) + assert table_axes({}, 1) is None + + with pytest.raises(ValueError, match="distinct valid one-based axes"): + table_axes( + {"x-nml-tools-ui": {"table": {"row-axis": 1, "column-axis": 1}}}, + 2, + ) + + +def test_parameter_array_display_round_trip_preserves_canonical_axis_order() -> None: + pytest.importorskip("numpy") + schema = { + "x-fortran-shape": [5, "n_units"], + "x-nml-tools-ui": { + "table": {"row-axis": 2, "column-axis": 1}, + }, + } + canonical = [ + [10.0, 11.0], + [20.0, 21.0], + [30.0, 31.0], + [0, 1], + [1, 0], + ] + + displayed = display_array(canonical, schema) + + assert displayed.tolist() == [ + [10.0, 20.0, 30.0, 0.0, 1.0], + [11.0, 21.0, 31.0, 1.0, 0.0], + ] + assert canonical_array(displayed, schema, rank=2) == canonical + + +def test_initial_array_broadcasts_parameter_vector_across_second_axis() -> None: + schema = {"x-fortran-shape": [5, "n_units"]} + + initialized = initial_array( + schema, + {"n_units": 2}, + [75.0, 200.0, 85.0, 1, 1], + 0.0, + ) + + assert initialized == [ + [75.0, 75.0], + [200.0, 200.0], + [85.0, 85.0], + [1, 1], + [1, 1], + ] + + +def test_derived_array_initialization_does_not_share_mutable_defaults() -> None: + initialized = initial_array( + {"x-fortran-shape": 2}, + {}, + None, + {"enabled": False}, + ) + + initialized[0]["enabled"] = True + assert initialized == [{"enabled": True}, {"enabled": False}] + + +def test_saved_array_shapes_are_preserved_or_rejected_without_coercion() -> None: + flexible = { + "x-fortran-shape": "max_items", + "x-fortran-flex-tail-dims": 1, + } + assert initial_array( + flexible, + {"max_items": 5}, + [10, 20], + 0, + strict=True, + ) == [10, 20] + + with pytest.raises(ValueError, match="does not match declared shape"): + initial_array( + {"x-fortran-shape": 5}, + {}, + [10, 20], + 0, + strict=True, + ) diff --git a/tests/test_gui_model.py b/tests/test_gui_model.py new file mode 100644 index 0000000..c669e37 --- /dev/null +++ b/tests/test_gui_model.py @@ -0,0 +1,266 @@ +"""Direct namelist persistence and profile selection checks.""" + +from pathlib import Path +from textwrap import dedent + +import pytest + +from nml_tools.gui.model import ( + create_virtual_project, + import_profile, + load_profile, + load_project, + overlay_values, + save_profiles, + suggestion, +) + + +def _write_project(root: Path, *, duplicate_output: bool = False) -> None: + schemas = root / "nml-schemas" + schemas.mkdir() + (schemas / "alpha.yml").write_text( + dedent( + """ + title: Alpha settings + x-fortran-namelist: alpha + type: object + properties: + count: + type: integer + label: + type: string + x-fortran-len: 32 + weights: + type: array + x-fortran-shape: n_items + items: + type: number + options: + type: object + x-fortran-type: options_t + properties: + enabled: + type: boolean + label: + type: string + x-fortran-len: 16 + settings: + type: array + x-fortran-shape: n_items + items: + type: object + x-fortran-type: setting_t + properties: + enabled: + type: boolean + name: + type: string + x-fortran-len: 16 + required: [count] + """ + ).lstrip(), + encoding="utf-8", + ) + + (schemas / "beta.yml").write_text( + dedent( + """ + title: Beta settings + x-fortran-namelist: beta + type: object + properties: + enabled: + type: boolean + """ + ).lstrip(), + encoding="utf-8", + ) + second_output = "main.nml" if duplicate_output else "secondary.nml" + (root / "nml-config.toml").write_text( + dedent( + f""" + [dimensions] + n_items = {{ default = 2 }} + + [[namelists]] + name = "alpha" + schema = "nml-schemas/alpha.yml" + + [[namelists]] + name = "beta" + schema = "nml-schemas/beta.yml" + + [[file_profiles]] + name = "secondary" + title = "Second profile" + default_file = "{second_output}" + namelists = ["beta"] + + [[file_profiles]] + name = "main" + default_file = "main.nml" + namelists = ["beta", "alpha"] + required = ["alpha"] + """ + ).lstrip(), + encoding="utf-8", + ) + + +def test_profiles_filter_in_toml_order_and_validate_names(tmp_path): + _write_project(tmp_path) + project = load_project(tmp_path, tmp_path / "output", {"MAIN": ["alpha"], "secondary": []}) + assert [p.name for p in project.profiles] == ["secondary", "main"] + assert [p.name for p in project.profile("main").pages] == ["alpha"] + assert [p.name for p in project.profile("secondary").pages] == ["beta"] + assert len(load_project(tmp_path, file_profiles={}).profiles) == 2 + for selection in ({"missing": []}, {"main": ["missing"]}, {"main": "alpha"}): + with pytest.raises(ValueError): + load_project(tmp_path, file_profiles=selection) + + +def test_save_reload_preserves_unselected_groups_and_uses_no_json(tmp_path): + _write_project(tmp_path) + project = load_project(tmp_path, file_profiles={"main": ["alpha"]}) + path = tmp_path / "main.nml" + untouched = "! keep this\n&beta enabled=.true. / ! keep too\n" + path.write_text(untouched + "&ALPHA count=1 /\n") + profile = project.profile("main") + values = { + "alpha": { + "count": 4, + "label": 'a "quote" / !', + "weights": [1.5, 2.5], + "options": {"enabled": True, "label": "single"}, + "settings": [{"enabled": True, "name": "first"}, {"enabled": False, "name": "second"}], + } + } + save_profiles(project, [(profile, values, {"n_items": 2})]) + saved = path.read_text() + assert saved.startswith(untouched) + assert "weights(2) = 2.5" in saved + assert "settings(1)%enabled = .true." in saved + assert load_profile(project, profile, {"n_items": 2}) == values + assert not list(tmp_path.glob("*.json")) + + +def test_all_saves_validate_before_replacing_any_file(tmp_path): + _write_project(tmp_path) + project = load_project(tmp_path) + path = tmp_path / "secondary.nml" + original = "&beta enabled=.true. /\n" + path.write_text(original) + with pytest.raises(ValueError): + save_profiles( + project, + [ + (project.profile("secondary"), {"beta": {"enabled": False}}, {}), + (project.profile("main"), {"alpha": {"count": "invalid"}}, {}), + ], + ) + assert path.read_text() == original + assert not (tmp_path / "main.nml").exists() + + +def test_missing_input_import_errors_and_partial_array_input(tmp_path): + _write_project(tmp_path) + project = load_project(tmp_path) + assert load_profile(project, project.profile("main"), {}) == {} + path = tmp_path / "imported.nml" + path.write_text("&alpha count=2 weights(2)=3.5 options%enabled=.true. /\n") + profile, values = import_profile(project, path, {}) + assert profile.default_file == "imported.nml" + assert values["alpha"]["weights"] == [0.0, 3.5] + assert values["alpha"]["options"] == {"enabled": True} + for text in ( + "&unknown a=1 /", + "&alpha count=2 / &alpha count=3 /", + "&alpha count=2 weights(3)=1 /", + "&alpha count=2", + ): + path.write_text(text) + with pytest.raises(ValueError): + import_profile(project, path, {}) + + +def test_array_defaults_use_fortran_order_and_declared_padding(): + schema = { + "type": "array", + "items": {"type": "integer"}, + "x-fortran-shape": [2, 2], + "default": [1, 2, 3, 4], + "examples": [[9]], + } + assert suggestion(schema, {}) == [[1, 3], [2, 4]] + assert suggestion({**schema, "x-fortran-default-order": "C"}, {}) == [[1, 2], [3, 4]] + assert suggestion({**schema, "default": [1], "x-fortran-default-pad": 7}, {}) == [ + [1, 7], + [7, 7], + ] + assert suggestion({"type": "integer", "default": 4, "examples": [9]}, {}) == 4 + + +def test_overlay_keeps_unspecified_derived_components(): + original = {"alpha": {"options": {"enabled": False, "label": "keep"}}} + merged = overlay_values(original, {"alpha": {"options": {"enabled": True}}}) + assert merged["alpha"]["options"] == {"enabled": True, "label": "keep"} + assert original["alpha"]["options"]["enabled"] is False + + +def test_multidimensional_and_deferred_arrays_round_trip(tmp_path): + _write_project(tmp_path) + project = load_project(tmp_path) + profile = project.profile("main") + alpha = next(page.schema for page in profile.pages if page.name == "alpha") + alpha["properties"]["weights"]["x-fortran-shape"] = [2, 2] + alpha["properties"]["settings"]["x-fortran-shape"] = [2, 2] + values = { + "alpha": { + "count": 1, + "weights": [[1.0, 2.0], [3.0, 4.0]], + "settings": [ + [{"enabled": True, "name": "a"}, {"enabled": False, "name": "b"}], + [{"enabled": False, "name": "c"}, {"enabled": True, "name": "d"}], + ], + } + } + save_profiles(project, [(profile, values, {})]) + assert "weights(2,1) = 3.0" in (tmp_path / "main.nml").read_text() + assert 'settings(1,2)%name = "b"' in (tmp_path / "main.nml").read_text() + assert load_profile(project, profile, {}) == values + alpha["properties"]["weights"]["x-fortran-shape"] = ":" + values["alpha"]["weights"] = [1.0, 2.0, 3.0] + save_profiles(project, [(profile, values, {})]) + assert load_profile(project, profile, {}) == values + + +def test_virtual_profiles_without_toml_profiles_and_output_safety(tmp_path): + _write_project(tmp_path) + config = tmp_path / "nml-config.toml" + config.write_text(config.read_text().split("[[file_profiles]]")[0]) + project = load_project(tmp_path) + assert not project.profiles + virtual = create_virtual_project(project, "custom", "custom.nml", ["alpha"]) + assert virtual.profiles[0].name == "custom" + with pytest.raises(ValueError, match="inside"): + create_virtual_project(project, "bad", "../outside.nml", ["alpha"]) + + +def test_save_rejects_strings_that_fortran_would_truncate(tmp_path): + _write_project(tmp_path) + project = load_project(tmp_path) + with pytest.raises(ValueError, match="length"): + save_profiles( + project, + [ + ( + project.profile("main"), + { + "alpha": {"count": 1, "label": "x" * 33}, + }, + {}, + ) + ], + ) + assert not (tmp_path / "main.nml").exists() diff --git a/tests/test_gui_widgets.py b/tests/test_gui_widgets.py new file mode 100644 index 0000000..4d67ff1 --- /dev/null +++ b/tests/test_gui_widgets.py @@ -0,0 +1,183 @@ +"""Offscreen checks of direct namelist editing and singleton fields.""" + +import os + +import pytest + +from nml_tools.gui.model import GuiProfile, GuiProject, NamelistPage + +os.environ.setdefault("QT_QPA_PLATFORM", "offscreen") +pytest.importorskip("qtpy") +try: + from qtpy.QtWidgets import QApplication, QMessageBox +except ImportError: + pytest.skip("Qt binding unavailable", allow_module_level=True) + +from nml_tools.gui.app import ConfigurationDialog, ProfileConfigTab +from nml_tools.gui.fields import ArrayField, FieldRow, ObjectField, ScalarField + + +@pytest.fixture(scope="module") +def application(): + return QApplication.instance() or QApplication([]) + + +@pytest.fixture +def project(tmp_path): + schema = { + "type": "object", + "x-fortran-namelist": "run", + "properties": { + "count": {"type": "integer", "default": 3, "examples": [99]}, + "label": {"type": "string", "x-fortran-len": 16, "default": "default"}, + "weights": { + "type": "array", + "x-fortran-shape": "n", + "items": {"type": "number", "default": 1.0}, + }, + "periods": { + "type": "array", + "x-fortran-shape": "n", + "items": { + "type": "object", + "x-fortran-type": "period_t", + "properties": {"year": {"type": "integer", "default": 2000}}, + }, + }, + }, + } + page = NamelistPage("run", "run", schema) + profile = GuiProfile("main", "main", "Main", None, "run.nml", (page,)) + return GuiProject(tmp_path, {}, {"n": 1}, (profile,), namelists=(page,)) + + +def test_singletons_keep_array_values_and_restore_schema_defaults(application): + schema = {"type": "array", "x-fortran-shape": "n", "items": {"type": "integer", "default": 7}} + row = FieldRow("counts", schema, [8], {"n": 1}) + assert isinstance(row.field, ArrayField) + assert isinstance(row.field.inline, ScalarField) + assert row.field.button.isHidden() + row.field.inline.set_value(12) + assert row.value() == [12] + row.reset({"n": 2}) + assert row.field.inline is None + assert row.value() == [7, 7] + row.reset({"n": 1}) + assert isinstance(row.field.inline, ScalarField) + assert row.value() == [7] + deferred = FieldRow("counts", {**schema, "x-fortran-shape": ":"}, [8], {}) + assert isinstance(deferred.field.inline, ScalarField) + assert not deferred.field.button.isHidden() + + +def test_derived_singletons_use_inline_object_fields(application, project): + schema = project.namelists[0].schema["properties"]["periods"] + row = FieldRow("periods", schema, [{"year": 2020}], {"n": 1}) + assert isinstance(row.field.inline, ObjectField) + row.field.inline.rows["year"].field.set_value(2025) + assert row.value() == [{"year": 2025}] + row.reset({"n": 1}) + assert row.value() == [{"year": 2000}] + scalar = ObjectField(schema["items"], {"year": 2021}, {}) + assert scalar.value() == {"year": 2021} + + +def test_dialog_load_overlay_save_reload_and_dimension_changes(application, project, monkeypatch): + errors = [] + monkeypatch.setattr(QMessageBox, "critical", lambda *args: errors.append(args[-1])) + path = project.root / "run.nml" + path.write_text('&run count=8 label="saved" weights(1)=2.0 periods(1)%year=2021 /\n') + dialog = ConfigurationDialog(project, initial_values={"main": {"run": {"count": 9}}}) + editor = dialog.editors[path] + assert editor.values()["run"]["count"] == 9 + assert editor.values()["run"]["label"] == "saved" + editor.forms["run"].rows["periods"].field.inline.rows["year"].field.set_value(2025) + dialog._save_all() + assert "periods(1)%year = 2025" in path.read_text() + dialog.config_tab.dimension_boxes["n"].setValue(2) + dialog.config_tab.run.click() + assert dialog.editors[path].forms["run"].rows["weights"].field.inline is None + dialog.config_tab.dimension_boxes["n"].setValue(1) + dialog.config_tab.run.click() + assert isinstance(dialog.editors[path].forms["run"].rows["weights"].field.inline, ScalarField) + reloaded = ConfigurationDialog(project) + assert reloaded.editors[path].values()["run"]["periods"] == [{"year": 2025}] + dialog.tabs.setCurrentWidget(dialog.plus_tab) + builder = dialog.tabs.currentWidget() + assert isinstance(builder, ProfileConfigTab) + builder.profile_name.setText("extra") + builder.default_filename.setText("extra.nml") + dialog._move_all(builder.available_schemas, builder.selected_schemas) + builder.run.click() + assert project.root / "extra.nml" in dialog.editors + dialog._save_all() + assert (project.root / "extra.nml").is_file() + assert not list(project.root.glob("*.json")) + assert not errors + dialog.close() + reloaded.close() + + +def test_invalid_existing_input_is_reported_and_never_replaced(application, project, monkeypatch): + errors = [] + monkeypatch.setattr(QMessageBox, "critical", lambda *args: errors.append(args[-1])) + path = project.root / "run.nml" + invalid = '&run count="wrong type" /' + path.write_text(invalid) + dialog = ConfigurationDialog(project) + assert errors and not dialog.editors + dialog._save_all() + assert path.read_text() == invalid + dialog.close() + + +def test_guidata_derived_edits_commit(application): + np = pytest.importorskip("numpy") + pytest.importorskip("guidata") + from guidata.widgets.arrayeditor import ArrayEditor + + from nml_tools.gui.fields import _derived_array_editor + + data = np.array([(2000,)], dtype=[("year", "i4")]).reshape(1, 1) + editor = _derived_array_editor(ArrayEditor, None) + try: + assert editor.setup_and_check(data) + editor._data.current_changes[("year", 0, 0)] = 2025 + editor.accept() + assert data["year"][0, 0] == 2025 + finally: + editor.close() + + +def test_imported_profile_can_choose_its_output_and_keep_loaded_values( + application, project, monkeypatch +): + errors = [] + monkeypatch.setattr(QMessageBox, "critical", lambda *args: errors.append(args[-1])) + path = project.root / "external.nml" + path.write_text('&run count=42 /\n') + dialog = ConfigurationDialog(project) + dialog.tabs.setCurrentWidget(dialog.plus_tab) + config = dialog.tabs.currentWidget() + config.source_combo.setCurrentIndex(config.source_combo.findData(str(path))) + config.profile_name.setText("copy") + config.default_filename.setText("copy.nml") + config.run.click() + copy_path = project.root / "copy.nml" + assert dialog.editors[copy_path].values()["run"]["count"] == 42 + dialog.editors[copy_path].save.click() + assert copy_path.exists() + assert path.read_text() == '&run count=42 /\n' + assert not errors + dialog.close() + + +def test_public_launch_forwards_profile_selection(application, monkeypatch, tmp_path): + from nml_tools.gui import app, launch_gui + + calls = [] + monkeypatch.setattr(app, "launch_gui", lambda *args: calls.append(args) or 0) + selected = {"main": ["run"]} + values = {"main": {"run": {"count": 8}}} + assert launch_gui(tmp_path, tmp_path / "out", selected, values, {"n": 1}) == 0 + assert calls == [(tmp_path, tmp_path / "out", selected, values, {"n": 1})] From 6ec0a33853ea9f37f71d84fb6bed0eebb1e2254f Mon Sep 17 00:00:00 2001 From: Sanjeev Bashyal Date: Sun, 13 Sep 2026 10:11:14 +0200 Subject: [PATCH 2/6] improved derived schema handling from reference schema --- src/nml_tools/gui/app.py | 70 ++++++++++++++++---- src/nml_tools/gui/arrays.py | 32 +++++---- src/nml_tools/gui/fields.py | 127 ++++++++++++++++++++++++++++++++++-- src/nml_tools/gui/model.py | 95 ++++++++++++++++++++++++--- tests/test_gui_arrays.py | 8 ++- tests/test_gui_model.py | 32 +++++++++ tests/test_gui_widgets.py | 98 ++++++++++++++++++++++++++-- 7 files changed, 419 insertions(+), 43 deletions(-) diff --git a/src/nml_tools/gui/app.py b/src/nml_tools/gui/app.py index 60b47cb..ba5fa8f 100644 --- a/src/nml_tools/gui/app.py +++ b/src/nml_tools/gui/app.py @@ -8,7 +8,7 @@ from pathlib import Path from typing import Any -from qtpy.QtCore import Qt +from qtpy.QtCore import QSignalBlocker, Qt from qtpy.QtWidgets import ( QAbstractItemView, QApplication, @@ -44,6 +44,7 @@ load_profile, load_project, overlay_values, + recover_dimensions, save_profiles, ) @@ -272,6 +273,12 @@ def __init__( self.dimensions = _normalize_dimensions( {} if initial_dimensions is None else initial_dimensions, project ) + self.dimension_overrides = initial_dimensions or {} + recovery_error = None + try: + self.dimensions = recover_dimensions(project, overrides=self.dimension_overrides) + except (OSError, ValueError) as exc: + recovery_error = str(exc) self.initial_values: dict[str, Any] = {} if initial_values is not None: if not isinstance(initial_values, Mapping): @@ -298,7 +305,7 @@ def __init__( root.addWidget(self.tabs, 1) self.plus_tab = QWidget(self.tabs) self.tabs.addTab(self.plus_tab, "+") - self.config_tab = self._add_config(primary=True) + self.config_tab: ProfileConfigTab | None = self._add_config(primary=True) for side in (QTabBar.LeftSide, QTabBar.RightSide): self.tabs.tabBar().setTabButton(self.tabs.indexOf(self.plus_tab), side, None) self.tabs.currentChanged.connect(self._tab_changed) @@ -314,7 +321,9 @@ def __init__( button.clicked.connect(callback) actions.addWidget(button) root.addLayout(actions) - if project.profiles: + if recovery_error: + QMessageBox.critical(self, "Invalid configuration", recovery_error) + elif project.profiles and self.config_tab is not None: self._run_configuration(self.config_tab) @staticmethod @@ -362,10 +371,13 @@ def _tab_changed(self, index: int) -> None: def _close_tab(self, index: int) -> None: widget = self.tabs.widget(index) - if widget is self.plus_tab: + if widget is None or widget is self.plus_tab: return try: - dirty = isinstance(widget, ProfileTab) and widget.values() != widget.saved_values + dirty = isinstance(widget, ProfileTab) and ( + widget.values() != widget.saved_values + or widget.dimensions != widget.saved_dimensions + ) except ValueError: dirty = True if dirty: @@ -375,9 +387,20 @@ def _close_tab(self, index: int) -> None: ): return self.editors = {path: tab for path, tab in self.editors.items() if tab is not widget} - self.config_tabs.discard(widget) - self.tabs.removeTab(index) - widget.deleteLater() + self._remove_tab(widget) + + def _remove_tab(self, widget: QWidget) -> None: + with QSignalBlocker(self.tabs): + index = self.tabs.indexOf(widget) + self.config_tabs.discard(widget) + self.tabs.removeTab(index) + if widget is self.config_tab: + self.config_tab = None + widget.deleteLater() + if self.tabs.count() == 1: + self.config_tab = self._add_config(primary=True) + elif self.tabs.currentWidget() is self.plus_tab: + self.tabs.setCurrentIndex(min(index, self.tabs.count() - 2)) def _browse(self, tab: ProfileConfigTab) -> None: name, _ = QFileDialog.getOpenFileName( @@ -393,13 +416,18 @@ def _browse(self, tab: ProfileConfigTab) -> None: def _select_source(self, tab: ProfileConfigTab) -> None: name = tab.source_combo.currentData() tab.source_path = Path(name) if name else None - if not tab.builder or not name: + if not name: return try: - profile, _ = import_profile(self.project, Path(name), tab.dimensions()) + dimensions = recover_dimensions(self.project, [Path(name)], self.dimension_overrides) + profile, _ = import_profile(self.project, Path(name), dimensions) except (OSError, ValueError) as exc: QMessageBox.critical(self, "Invalid namelist", str(exc)) return + for key, size in dimensions.items(): + tab.dimension_boxes[key].setValue(size) + if not tab.builder: + return tab.profile_name.setText(profile.name) tab.default_filename.setText(profile.default_file) tab.available_schemas.clear() @@ -411,6 +439,7 @@ def _select_source(self, tab: ProfileConfigTab) -> None: def _run_configuration(self, config: ProfileConfigTab) -> None: prepared: list[ProfileTab] = [] + allow_shrink = False try: dimensions = _normalize_dimensions(config.dimensions(), self.project) imported: dict[str, Any] | None = None @@ -435,6 +464,23 @@ def _run_configuration(self, config: ProfileConfigTab) -> None: previous = self.editors.get(target) if previous is not None and (config.builder or config.source_path is not None): raise ValueError(f"'{target.name}' is already open") + if ( + previous + and not allow_shrink + and any(size < previous.dimensions[key] for key, size in dimensions.items()) + ): + if ( + QMessageBox.question( + self, + "Resize arrays", + "Smaller dimensions will discard trailing array values. Continue?", + ) + != QMessageBox.Yes + ): + for editor in prepared: + editor.deleteLater() + return + allow_shrink = True values = ( previous.values() if previous @@ -466,9 +512,7 @@ def _run_configuration(self, config: ProfileConfigTab) -> None: self._put_editor(editor) self.dimensions = dimensions if config.builder or config.source_path is not None: - self.config_tabs.discard(config) - self.tabs.removeTab(self.tabs.indexOf(config)) - config.deleteLater() + self._remove_tab(config) except (OSError, ValueError, KeyError) as exc: for editor in prepared: editor.deleteLater() diff --git a/src/nml_tools/gui/arrays.py b/src/nml_tools/gui/arrays.py index 798e6fa..0387930 100644 --- a/src/nml_tools/gui/arrays.py +++ b/src/nml_tools/gui/arrays.py @@ -178,10 +178,16 @@ def initial_array( leaf_default: Any, *, strict: bool = False, + resize: bool = False, + defaults: list[Any] | None = None, ) -> list[Any]: """Fit a saved/default/example value to the resolved canonical shape.""" shape = resolve_shape(schema, sizes, value) - result = cast(list[Any], _filled(shape, leaf_default)) + result = ( + copy.deepcopy(defaults) + if defaults is not None + else cast(list[Any], _filled(shape, leaf_default)) + ) if value is None: return result if not isinstance(value, list): @@ -191,17 +197,19 @@ def initial_array( current_shape = _list_shape(value) if strict: validate_array_shape(schema, sizes, value) - return copy.deepcopy(value) - if current_shape and flex_tail_dims(schema, len(shape)): - try: - validate_array_shape(schema, sizes, value) - except ValueError: - pass - else: - return copy.deepcopy(value) - if current_shape == shape: - return copy.deepcopy(value) - + if strict or resize or (current_shape and len(current_shape) == len(shape)): + if len(current_shape) != len(shape): + raise ValueError("saved array rank does not match its declared shape") + + def overlay(target: list[Any], source: list[Any]) -> None: + for index, item in enumerate(source[: len(target)]): + if isinstance(target[index], list): + overlay(target[index], item) + else: + target[index] = copy.deepcopy(item) + + overlay(result, value) + return result flat = list(_flatten(value)) if not flat: return result diff --git a/src/nml_tools/gui/fields.py b/src/nml_tools/gui/fields.py index 57bf14a..2a537a4 100644 --- a/src/nml_tools/gui/fields.py +++ b/src/nml_tools/gui/fields.py @@ -5,6 +5,7 @@ import copy import math from collections.abc import Mapping +from itertools import product from typing import Any, cast from qtpy.QtWidgets import ( @@ -13,19 +14,22 @@ QFormLayout, QGroupBox, QHBoxLayout, + QHeaderView, QLabel, QLineEdit, QMessageBox, QPushButton, + QTableWidget, + QTableWidgetItem, QWidget, ) +from ..schema import DERIVED_REF_ORIGIN_KEY from .arrays import ( array_shape, axis_labels, canonical_array, display_array, - flex_tail_dims, initial_array, resolve_shape, table_axes, @@ -184,6 +188,12 @@ def __init__( candidate, suggestion(items, sizes), strict=saved and not fit_existing, + resize=saved and fit_existing, + defaults=suggestion( + {**schema, "x-fortran-shape": list(resolve_shape(schema, sizes, candidate))}, sizes + ) + if saved and fit_existing + else None, ) layout = QHBoxLayout(self) layout.setContentsMargins(0, 0, 0, 0) @@ -216,7 +226,6 @@ def _update_summary(self) -> None: self.layout().insertWidget(0, self.inline, 1) raw = self.schema.get("x-fortran-shape") resizable = raw == ":" or (isinstance(raw, list) and ":" in raw) - resizable = resizable or flex_tail_dims(self.schema, len(shape)) > 0 self.summary.setVisible(self.inline is None) self.button.setVisible(self.inline is None or resizable) self.button.setText("Resize…" if self.inline is not None else "Edit array…") @@ -241,7 +250,7 @@ def _edit(self) -> None: str(self.schema.get("title", self.name)), xlabels=xlabels, ylabels=ylabels, - variable_size=flex_tail_dims(self.schema, rank) > 0 or deferred, + variable_size=deferred, ): return if _exec(editor) != _accepted(editor): @@ -368,6 +377,93 @@ def set_value(self, value: Any, sizes: Mapping[str, int]) -> None: self.field = replacement +class DerivedTable(QTableWidget): + """Same-reference objects as rows, reusing the existing component editors.""" + + def __init__( + self, + schemas: Mapping[str, Mapping[str, Any]], + values: Mapping[str, Any], + sizes: Mapping[str, int], + required: set[str], + parent: QWidget, + *, + fit_arrays: bool, + ) -> None: + super().__init__(parent) + self.schemas, self.sizes = schemas, sizes + self.data: dict[str, Any] = {} + self.objects: dict[tuple[str, tuple[int, ...]], ObjectField] = {} + first = next(iter(schemas.values())) + first = first["items"] if first["type"] == "array" else first + columns = list(first["properties"]) + self.setColumnCount(len(columns)) + self.setHorizontalHeaderLabels(columns) + header = cast(QHeaderView, self.horizontalHeader()) + header.setSectionResizeMode(QHeaderView.ResizeToContents) + for name, schema in schemas.items(): + value = values.get(name, MISSING) + is_array = schema["type"] == "array" + item = schema["items"] if is_array else schema + if is_array: + self.data[name] = initial_array( + schema, + sizes, + suggestion(schema, sizes) if value is MISSING else value, + suggestion(item, sizes), + strict=value is not MISSING and not fit_arrays, + resize=value is not MISSING and fit_arrays, + defaults=suggestion( + {**schema, "x-fortran-shape": list(resolve_shape(schema, sizes, value))}, + sizes, + ) + if value is not MISSING and fit_arrays + else None, + ) + else: + self.data[name] = {} if value is MISSING else value + shape = array_shape(self.data[name]) if is_array else () + for indices in product(*(range(n) for n in shape)): + obj = ObjectField(item, _nested_get(self.data[name], indices), sizes, self) + obj.hide() + row = self.rowCount() + self.insertRow(row) + suffix = "(" + ",".join(str(i + 1) for i in indices) + ")" if indices else "" + label = QTableWidgetItem(name + suffix + (" *" if name.lower() in required else "")) + label.setToolTip(str(schema.get("description", schema.get("title", name)))) + self.setVerticalHeaderItem(row, label) + for column, component in enumerate(columns): + self.setCellWidget(row, column, obj.rows[component]) + label = QTableWidgetItem( + component + (" *" if component in item.get("required", []) else "") + ) + label.setToolTip( + str(item["properties"][component].get("description", component)) + ) + self.setHorizontalHeaderItem(column, label) + self.objects[name, indices] = obj + self.resizeRowsToContents() + self.setMinimumHeight( + min(400, header.height() + sum(self.rowHeight(i) for i in range(self.rowCount())) + 4) + ) + + def values(self) -> dict[str, Any]: + result = copy.deepcopy(self.data) + for (name, indices), obj in self.objects.items(): + if indices: + _nested_get(result[name], indices[:-1])[indices[-1]] = obj.value() + else: + result[name] = obj.value() + return result + + def reset(self) -> None: + defaults = {name: suggestion(schema, self.sizes) for name, schema in self.schemas.items()} + for (name, indices), obj in self.objects.items(): + obj.reset(self.sizes) + for component, value in _nested_get(defaults[name], indices).items(): + obj.rows[component].set_value(value, self.sizes) + + class NamelistForm(QWidget): """Editable form for one namelist schema.""" @@ -390,10 +486,29 @@ def __init__( required = {item.lower() for item in schema.get("required", []) if isinstance(item, str)} layout = QFormLayout(self) self.rows: dict[str, FieldRow] = {} + self.tables: list[DerivedTable] = [] + groups: dict[tuple[str, ...], dict[str, Mapping[str, Any]]] = {} + identities = {} + for name, child in properties.items(): + item = child.get("items", {}) if child.get("type") == "array" else child + origin = item.get(DERIVED_REF_ORIGIN_KEY) + if item.get("type") == "object" and origin: + identity = tuple(origin["identity"]) + identities[name] = identity + groups.setdefault(identity, {})[name] = child for name, child in properties.items(): if not isinstance(name, str) or not isinstance(child, Mapping): continue is_required = name.lower() in required + children = groups.get(identities.get(name, ())) + if children: + if name == next(iter(children)): + table = DerivedTable( + children, source, sizes, required, self, fit_arrays=fit_arrays + ) + layout.addRow(table) + self.tables.append(table) + continue row = FieldRow( name, child, @@ -411,11 +526,15 @@ def values(self) -> dict[str, Any]: value = row.value() if value is not MISSING: result[name] = value - return result + for table in self.tables: + result.update(table.values()) + return {name: result[name] for name in self.schema["properties"] if name in result} def reset(self) -> None: for row in self.rows.values(): row.reset(self.sizes) + for table in self.tables: + table.reset() def _field_widget( diff --git a/src/nml_tools/gui/model.py b/src/nml_tools/gui/model.py index cb89863..3ff5c23 100644 --- a/src/nml_tools/gui/model.py +++ b/src/nml_tools/gui/model.py @@ -10,12 +10,12 @@ from dataclasses import dataclass, replace from itertools import product from pathlib import Path -from typing import Any, Mapping +from typing import Any, Mapping, cast import click -from .._namelist_eval import EvaluatedGroup, LeafState, evaluate_group -from .._namelist_parser import parse_namelist +from .._namelist_eval import EvaluatedGroup, LeafState, _expand_values, _value_count, evaluate_group +from .._namelist_parser import RawValue, ScalarSelector, parse_namelist from ..cli import ( _iter_file_profiles, _load_config_checked, @@ -27,7 +27,7 @@ from ..codegen_fortran import _format_scalar_default from ..schema import SchemaResolver from ..validate import _scalar_constraints, _validate_scalar_value, validate_schema_defaults -from .arrays import flex_tail_dims, initial_array, resolve_shape, validate_array_shape +from .arrays import initial_array, resolve_shape MISSING = object() @@ -332,16 +332,15 @@ def _evaluated_array( if not isinstance(items, Mapping): raise ValueError("array field must define object 'items'") shape = list(resolve_shape(schema, sizes)) - flexible = flex_tail_dims(schema, len(shape)) raw = schema["x-fortran-shape"] raw = raw if isinstance(raw, list) else [raw] for axis in range(len(shape)): - if raw[axis] != ":" and axis < len(shape) - flexible: + if raw[axis] != ":": continue used = [coordinates[axis] for coordinates, _, _ in states if coordinates] if used: shape[axis] = max(used) - result = suggestion({**schema, "x-fortran-shape": shape}, sizes) + result = cast(list[Any], suggestion({**schema, "x-fortran-shape": shape}, sizes)) components = _component_names(items) if items.get("type") == "object" else {} for coordinates, component, state in states: target = result @@ -437,10 +436,10 @@ def _normalize_value( if kind == "array": if not isinstance(value, list): raise ValueError(f"'{path}' must be an array") - validate_array_shape(schema, sizes, value) items = schema.get("items") if not isinstance(items, Mapping): raise ValueError(f"array '{path}' must define object items") + value = initial_array(schema, sizes, value, suggestion(items, sizes), strict=True) def normalize_items(node: Any, indices: tuple[int, ...] = ()) -> Any: if isinstance(node, list): @@ -538,6 +537,86 @@ def _parsed_groups(path: Path) -> tuple[str, dict[str, Any]]: return text, groups +def recover_dimensions( + project: GuiProject, + paths: Iterable[Path] | None = None, + overrides: Mapping[str, int] | None = None, +) -> dict[str, int]: + """Infer editable sizes from saved indices before evaluating the namelists.""" + if paths is None: + paths = (_profile_path(project, profile) for profile in project.profiles) + pages = project.namelists or tuple(page for p in project.profiles for page in p.pages) + schemas = {page.key: page.schema for page in pages} + extents: dict[str, int] = {} + explicit: dict[str, int] = {} + for path in dict.fromkeys(paths): + if not path.exists(): + continue + _, groups = _parsed_groups(path) + saved: dict[str, int] = {} + for key, group in groups.items(): + properties = { + name.lower(): prop + for name, prop in schemas.get(key, {}).get("properties", {}).items() + } + for assignment in group.assignments: + part = assignment.designator.parts[0] + name = part.name.lower() + prop = properties.get(name, {}) + if ( + name in project.default_dimensions + and prop.get("type") == "integer" + and len(assignment.designator.parts) == 1 + and not part.selectors + ): + for value in _expand_values(assignment): + if isinstance(value, RawValue) and not value.quoted: + saved[name] = int(value.source_text.split("_")[0]) + if prop.get("type") != "array": + continue + raw = prop["x-fortran-shape"] + shape = raw if isinstance(raw, list) else [raw] + selectors = part.selectors[0].selectors if part.selectors else () + count = _value_count(assignment) + items = prop["items"] + if items.get("type") == "object" and len(assignment.designator.parts) == 1: + count = (count + len(items["properties"]) - 1) // len(items["properties"]) + for axis, token in enumerate(shape): + dimension = str(token).lower() + if dimension not in project.default_dimensions: + continue + selector = selectors[axis] if axis < len(selectors) else None + extent = 0 + if isinstance(selector, ScalarSelector): + extent = selector.value + (max(0, count - 1) if len(shape) == 1 else 0) + elif selector is not None and selector.upper is not None: + stride = selector.stride or 1 + indices = range( + selector.lower or 1, selector.upper + (1 if stride > 0 else -1), stride + ) + extent = max(indices[0], indices[-1]) if indices else 0 + elif len(shape) == 1: + lower = (selector.lower or 1) if selector is not None else 1 + stride = (selector.stride or 1) if selector is not None else 1 + extent = max(lower, lower + (count - 1) * stride) if count else 0 + if extent > 0: + extents[dimension] = max(extents.get(dimension, 0), extent) + for name, size in saved.items(): + if name in explicit and explicit[name] != size: + raise ValueError(f"conflicting saved dimension '{name}'") + explicit[name] = size + _normalize_dimensions(overrides or {}, project) + # ponytail: sparse external files give lower bounds, not declared capacities. + dimensions = _normalize_dimensions( + {**extents, **explicit, **{name.lower(): size for name, size in (overrides or {}).items()}}, + project, + ) + for name, extent in extents.items(): + if dimensions[name] < extent: + raise ValueError(f"dimension '{name}' is smaller than saved extent {extent}") + return dimensions + + def load_profile( project: GuiProject, profile: GuiProfile, diff --git a/tests/test_gui_arrays.py b/tests/test_gui_arrays.py index eef955f..0c35adb 100644 --- a/tests/test_gui_arrays.py +++ b/tests/test_gui_arrays.py @@ -118,7 +118,7 @@ def test_derived_array_initialization_does_not_share_mutable_defaults() -> None: assert initialized == [{"enabled": True}, {"enabled": False}] -def test_saved_array_shapes_are_preserved_or_rejected_without_coercion() -> None: +def test_saved_arrays_are_completed_without_repeating_or_moving_values() -> None: flexible = { "x-fortran-shape": "max_items", "x-fortran-flex-tail-dims": 1, @@ -129,7 +129,11 @@ def test_saved_array_shapes_are_preserved_or_rejected_without_coercion() -> None [10, 20], 0, strict=True, - ) == [10, 20] + ) == [10, 20, 0, 0, 0] + assert initial_array({"x-fortran-shape": [2, 3]}, {}, [[1, 2], [3, 4]], 9, resize=True) == [ + [1, 2, 9], + [3, 4, 9], + ] with pytest.raises(ValueError, match="does not match declared shape"): initial_array( diff --git a/tests/test_gui_model.py b/tests/test_gui_model.py index c669e37..c5f7555 100644 --- a/tests/test_gui_model.py +++ b/tests/test_gui_model.py @@ -11,6 +11,7 @@ load_profile, load_project, overlay_values, + recover_dimensions, save_profiles, suggestion, ) @@ -184,6 +185,37 @@ def test_missing_input_import_errors_and_partial_array_input(tmp_path): import_profile(project, path, {}) +def test_complete_arrays_and_recover_dimensions_before_loading(tmp_path): + _write_project(tmp_path) + project = load_project(tmp_path) + profile = project.profile("main") + schema = next(page.schema for page in project.namelists if page.key == "alpha") + weights = schema["properties"]["weights"] + weights.update({"x-fortran-shape": [2, "n_items"], "x-fortran-flex-tail-dims": 1}) + path = tmp_path / "main.nml" + for size in (3, 1): + values = {"alpha": {"count": 1, "weights": [[4.0], [5.0]]}} + save_profiles(project, [(profile, values, {"n_items": size})]) + assert f"weights(2,{size}) =" in path.read_text() + dimensions = recover_dimensions(project) + assert dimensions == {"n_items": size} + assert load_profile(project, profile, dimensions)["alpha"]["weights"] == [ + [4.0] + [0.0] * (size - 1), + [5.0] + [0.0] * (size - 1), + ] + weights["x-fortran-shape"] = "n_items" + path.write_text("&alpha count=1 weights(2:6:2)=3*1.0 /\n") + assert recover_dimensions(project) == {"n_items": 6} + assert recover_dimensions(project, overrides={"N_ITEMS": 7}) == {"n_items": 7} + with pytest.raises(ValueError, match="smaller than saved extent"): + recover_dimensions(project, overrides={"n_items": 1}) + path.write_text("&alpha count=1 weights(2:)=3*1.0 /\n") + assert recover_dimensions(project) == {"n_items": 4} + schema["properties"]["n_items"] = {"type": "integer"} + path.write_text("&alpha count=1 n_items=7 weights(1)=4.0 /\n") + assert recover_dimensions(project) == {"n_items": 7} + + def test_array_defaults_use_fortran_order_and_declared_padding(): schema = { "type": "array", diff --git a/tests/test_gui_widgets.py b/tests/test_gui_widgets.py index 4d67ff1..d67e96b 100644 --- a/tests/test_gui_widgets.py +++ b/tests/test_gui_widgets.py @@ -5,6 +5,7 @@ import pytest from nml_tools.gui.model import GuiProfile, GuiProject, NamelistPage +from nml_tools.schema import resolve_schema os.environ.setdefault("QT_QPA_PLATFORM", "offscreen") pytest.importorskip("qtpy") @@ -14,7 +15,7 @@ pytest.skip("Qt binding unavailable", allow_module_level=True) from nml_tools.gui.app import ConfigurationDialog, ProfileConfigTab -from nml_tools.gui.fields import ArrayField, FieldRow, ObjectField, ScalarField +from nml_tools.gui.fields import ArrayField, FieldRow, NamelistForm, ObjectField, ScalarField @pytest.fixture(scope="module") @@ -33,7 +34,9 @@ def project(tmp_path): "weights": { "type": "array", "x-fortran-shape": "n", - "items": {"type": "number", "default": 1.0}, + "items": {"type": "number"}, + "default": [1.0], + "x-fortran-default-repeat": True, }, "periods": { "type": "array", @@ -85,6 +88,7 @@ def test_derived_singletons_use_inline_object_fields(application, project): def test_dialog_load_overlay_save_reload_and_dimension_changes(application, project, monkeypatch): errors = [] monkeypatch.setattr(QMessageBox, "critical", lambda *args: errors.append(args[-1])) + monkeypatch.setattr(QMessageBox, "question", lambda *args: QMessageBox.Yes) path = project.root / "run.nml" path.write_text('&run count=8 label="saved" weights(1)=2.0 periods(1)%year=2021 /\n') dialog = ConfigurationDialog(project, initial_values={"main": {"run": {"count": 9}}}) @@ -97,9 +101,15 @@ def test_dialog_load_overlay_save_reload_and_dimension_changes(application, proj dialog.config_tab.dimension_boxes["n"].setValue(2) dialog.config_tab.run.click() assert dialog.editors[path].forms["run"].rows["weights"].field.inline is None + assert dialog.editors[path].values()["run"]["weights"] == [2.0, 1.0] + dialog._save_all() + loaded = ConfigurationDialog(project) + assert loaded.config_tab.dimension_boxes["n"].value() == 2 + loaded.close() dialog.config_tab.dimension_boxes["n"].setValue(1) dialog.config_tab.run.click() assert isinstance(dialog.editors[path].forms["run"].rows["weights"].field.inline, ScalarField) + dialog._save_all() reloaded = ConfigurationDialog(project) assert reloaded.editors[path].values()["run"]["periods"] == [{"year": 2025}] dialog.tabs.setCurrentWidget(dialog.plus_tab) @@ -155,11 +165,12 @@ def test_imported_profile_can_choose_its_output_and_keep_loaded_values( errors = [] monkeypatch.setattr(QMessageBox, "critical", lambda *args: errors.append(args[-1])) path = project.root / "external.nml" - path.write_text('&run count=42 /\n') + path.write_text("&run count=42 weights(3)=9.0 /\n") dialog = ConfigurationDialog(project) dialog.tabs.setCurrentWidget(dialog.plus_tab) config = dialog.tabs.currentWidget() config.source_combo.setCurrentIndex(config.source_combo.findData(str(path))) + assert config.dimension_boxes["n"].value() == 3 config.profile_name.setText("copy") config.default_filename.setText("copy.nml") config.run.click() @@ -167,7 +178,7 @@ def test_imported_profile_can_choose_its_output_and_keep_loaded_values( assert dialog.editors[copy_path].values()["run"]["count"] == 42 dialog.editors[copy_path].save.click() assert copy_path.exists() - assert path.read_text() == '&run count=42 /\n' + assert path.read_text() == "&run count=42 weights(3)=9.0 /\n" assert not errors dialog.close() @@ -181,3 +192,82 @@ def test_public_launch_forwards_profile_selection(application, monkeypatch, tmp_ values = {"main": {"run": {"count": 8}}} assert launch_gui(tmp_path, tmp_path / "out", selected, values, {"n": 1}) == 0 assert calls == [(tmp_path, tmp_path / "out", selected, values, {"n": 1})] + + +def test_active_and_last_config_tabs_close(application, project): + dialog = ConfigurationDialog(project) + dialog.tabs.setCurrentWidget(dialog.plus_tab) + builder = dialog.tabs.currentWidget() + dialog._close_tab(dialog.tabs.indexOf(builder)) + assert dialog.tabs.indexOf(builder) == -1 and len(dialog.config_tabs) == 1 + dialog._close_tab(dialog.tabs.indexOf(dialog.config_tab)) + assert dialog.config_tab is None + dialog._close_tab(dialog.tabs.indexOf(next(iter(dialog.editors.values())))) + last = dialog.config_tab + dialog._close_tab(dialog.tabs.indexOf(last)) + assert dialog.config_tab is not last and dialog.tabs.count() == 2 + dialog.close() + + +def test_shared_reference_table_edit_reset_and_round_trip(application, tmp_path): + from nml_tools.gui.model import load_profile, recover_dimensions, save_profiles + + schema = resolve_schema( + { + "type": "object", + "x-fortran-namelist": "run", + "$defs": { + "period": { + "type": "object", + "x-fortran-type": "period_t", + "properties": { + "year": {"type": "integer", "default": 2000}, + "enabled": {"type": "boolean", "default": False}, + }, + "required": ["year"], + } + }, + "properties": { + "start": {"$ref": "#/$defs/period", "default": {"year": 2001}}, + "stop": {"$ref": "#/$defs/period", "default": {"year": 2002}}, + "periods": { + "type": "array", + "x-fortran-shape": "n", + "items": {"$ref": "#/$defs/period"}, + }, + }, + } + ) + form = NamelistForm(schema, None, {"n": 2}) + (table,) = form.tables + form.show() + application.processEvents() + assert table.cellWidget(0, 0).isVisible() + assert (table.rowCount(), table.columnCount()) == (4, 2) + assert table.horizontalHeaderItem(0).text() == "year *" + table.objects["periods", (1,)].rows["year"].field.set_value(2025) + table.objects["periods", (1,)].rows["enabled"].field.set_value(True) + values = form.values() + assert values["periods"][1] == {"year": 2025, "enabled": True} + page = NamelistPage("run", "run", schema) + profile = GuiProfile("main", "main", "Main", None, "run.nml", (page,)) + project = GuiProject(tmp_path, {}, {"n": 100}, (profile,), namelists=(page,)) + save_profiles(project, [(profile, {"run": values}, {"n": 2})]) + assert load_profile(project, profile, recover_dimensions(project)) == {"run": values} + form.reset() + assert form.values()["start"]["year"] == 2001 + assert form.values()["stop"]["year"] == 2002 + assert form.values()["periods"][1]["year"] == 2000 + for name, indices in (("start", ()), ("periods", (0,))): + single = NamelistForm( + {**schema, "properties": {name: schema["properties"][name]}}, None, {"n": 1} + ) + (single_table,) = single.tables + assert (single_table.rowCount(), single_table.columnCount()) == (1, 2) + assert single_table.horizontalHeaderItem(0).text() == "year *" + defaults = single.values() + single_table.objects[name, indices].rows["year"].field.set_value(2035) + edited = {"year": 2035, "enabled": False} + assert single.values()[name] == ([edited] if indices else edited) + single.reset() + assert single.values() == defaults From fae1800a0230d10b643c6a3f9a5264bb83465ed9 Mon Sep 17 00:00:00 2001 From: Sanjeev Bashyal Date: Tue, 22 Sep 2026 13:47:07 +0200 Subject: [PATCH 3/6] Fix: added a modified tracking to not write empty namelist of format: file-path to prevent root folder assumption by mhmv6 --- src/nml_tools/gui/fields.py | 244 ++++++++++++++++++++++++++++++++---- src/nml_tools/gui/model.py | 21 +++- 2 files changed, 240 insertions(+), 25 deletions(-) diff --git a/src/nml_tools/gui/fields.py b/src/nml_tools/gui/fields.py index 2a537a4..5dde006 100644 --- a/src/nml_tools/gui/fields.py +++ b/src/nml_tools/gui/fields.py @@ -34,7 +34,7 @@ resolve_shape, table_axes, ) -from .model import MISSING, overlay_values, suggestion +from .model import MISSING, InputArray, overlay_values, suggestion def _exec(dialog: Any) -> int: @@ -62,6 +62,17 @@ def accept(self) -> None: return DerivedArrayEditor(parent) +def _seeded(schema: Mapping[str, Any]) -> bool: + if "default" in schema or schema.get("examples"): + return True + kind = schema.get("type") + if kind == "array": + return _seeded(schema["items"]) + if kind == "object": + return any(_seeded(child) for child in schema["properties"].values()) + return False + + class ScalarField(QWidget): def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None = None): super().__init__(parent) @@ -83,6 +94,13 @@ def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None self.control = control layout.addWidget(control) self.set_value(value) + self.modified = False + signal = ( + control.textEdited if isinstance(control, QLineEdit) + else control.toggled if isinstance(control, QCheckBox) + else control.currentIndexChanged + ) + signal.connect(lambda *_: setattr(self, "modified", True)) def set_value(self, value: Any) -> None: if isinstance(self.control, QComboBox): @@ -92,6 +110,7 @@ def set_value(self, value: Any) -> None: self.control.setChecked(bool(value)) else: self.control.setText(str(value)) + self.modified = True def value(self) -> Any: if isinstance(self.control, QComboBox): @@ -114,6 +133,7 @@ def value(self) -> Any: def reset(self, sizes: Mapping[str, int]) -> None: self.set_value(suggestion(self.schema, sizes)) + self.modified = False class ObjectField(QGroupBox): @@ -132,7 +152,9 @@ def __init__( properties = schema.get("properties") if not isinstance(properties, Mapping): raise ValueError("derived field must define object 'properties'") - source = suggestion(schema, sizes) + source = ( + suggestion(schema, sizes) if "default" in schema or schema.get("examples") else {} + ) if isinstance(value, Mapping): source = overlay_values(source, value) required = {item.lower() for item in schema.get("required", []) if isinstance(item, str)} @@ -156,9 +178,16 @@ def value(self) -> dict[str, Any]: return result def reset(self, sizes: Mapping[str, int]) -> None: - defaults = suggestion(self.schema, sizes) + defaults = ( + suggestion(self.schema, sizes) + if "default" in self.schema or self.schema.get("examples") + else {} + ) for name, row in self.rows.items(): - row.set_value(defaults[name], sizes) + if name in defaults: + row.set_value(defaults[name], sizes) + else: + row.reset(sizes) class ArrayField(QWidget): @@ -195,6 +224,14 @@ def __init__( if saved and fit_existing else None, ) + shape = array_shape(self._value) + positions = set(product(*(range(1, size + 1) for size in shape))) + self.assigned = ( + set(value.assigned) if isinstance(value, InputArray) + else positions if saved else set() + ) & positions + if _seeded(schema): + self.assigned.update(positions) layout = QHBoxLayout(self) layout.setContentsMargins(0, 0, 0, 0) self.summary = QLabel(self) @@ -207,12 +244,20 @@ def __init__( def value(self) -> list[Any]: if self.inline is not None: - return [self.inline.value()] - return copy.deepcopy(self._value) + value = self.inline.value() + if getattr(self.inline, "modified", False) or isinstance(value, dict) and value: + self.assigned.add((1,)) + return InputArray([value], set(self.assigned)) + return InputArray(copy.deepcopy(self._value), set(self.assigned)) def reset(self, sizes: Mapping[str, int]) -> None: self.sizes = sizes self._value = suggestion(self.schema, sizes) + shape = array_shape(self._value) + self.assigned = ( + set(product(*(range(1, size + 1) for size in shape))) + if _seeded(self.schema) else set() + ) self._update_summary() def _update_summary(self) -> None: @@ -236,7 +281,8 @@ def _edit(self) -> None: import numpy as np from guidata.widgets.arrayeditor import ArrayEditor # type: ignore[import-untyped] - self._value = self.value() + before = self.value() + self._value = before rank = len(resolve_shape(self.schema, self.sizes, self._value)) derived = self.items.get("type") == "object" canonical = self._structured_array(np) if derived else self._intrinsic_array(np) @@ -260,6 +306,19 @@ def _edit(self) -> None: self._value = self._objects_from_structured(edited, rank, np) else: self._value = canonical_array(edited, self.schema, rank) + shape = array_shape(self._value) + old_shape = array_shape(before) + self.assigned = { + indices + for indices in self.assigned + if all(index <= size for index, size in zip(indices, shape)) + } + for indices in product(*(range(size) for size in shape)): + if ( + any(index >= size for index, size in zip(indices, old_shape)) + or _nested_get(before, indices) != _nested_get(self._value, indices) + ): + self.assigned.add(tuple(index + 1 for index in indices)) self._update_summary() except (ImportError, RuntimeError, TypeError, ValueError) as exc: QMessageBox.critical(self, "Array editor", str(exc)) @@ -351,6 +410,7 @@ def __init__( self.name = name self.schema = schema self.sizes = sizes + self._provided = value is not MISSING layout = QHBoxLayout(self) layout.setContentsMargins(0, 0, 0, 0) initial = value @@ -363,14 +423,25 @@ def __init__( layout.addWidget(self.field, 1) def value(self) -> Any: - return self.field.value() + if isinstance(self.field, ScalarField) and not ( + self._provided or _seeded(self.schema) or self.field.modified + ): + return MISSING + value = self.field.value() + if isinstance(value, InputArray) and not value.assigned: + return MISSING + return MISSING if isinstance(self.field, ObjectField) and not value else value def reset(self, sizes: Mapping[str, int]) -> None: - self.set_value(suggestion(self.schema, sizes), sizes) + self.set_value(MISSING, sizes) def set_value(self, value: Any, sizes: Mapping[str, int]) -> None: + self._provided = value is not MISSING self.sizes = sizes - replacement = _field_widget(self.name, self.schema, value, sizes, self) + initial = value + if value is MISSING and self.schema.get("type") not in {"array", "object"}: + initial = suggestion(self.schema, sizes) + replacement = _field_widget(self.name, self.schema, initial, sizes, self) replacement.setToolTip(self.field.toolTip()) self.layout().replaceWidget(self.field, replacement) self.field.deleteLater() @@ -406,7 +477,7 @@ def __init__( is_array = schema["type"] == "array" item = schema["items"] if is_array else schema if is_array: - self.data[name] = initial_array( + dense = initial_array( schema, sizes, suggestion(schema, sizes) if value is MISSING else value, @@ -420,11 +491,25 @@ def __init__( if value is not MISSING and fit_arrays else None, ) + shape = array_shape(dense) + positions = set(product(*(range(1, n + 1) for n in shape))) + source_positions = ( + set(value.assigned) if isinstance(value, InputArray) + else positions if value is not MISSING else set() + ) & positions + self.data[name] = InputArray( + dense, source_positions | (positions if _seeded(schema) else set()) + ) else: self.data[name] = {} if value is MISSING else value shape = array_shape(self.data[name]) if is_array else () for indices in product(*(range(n) for n in shape)): - obj = ObjectField(item, _nested_get(self.data[name], indices), sizes, self) + source = ( + _nested_get(self.data[name], indices) + if not is_array or tuple(index + 1 for index in indices) in source_positions + else MISSING + ) + obj = ObjectField(item, source, sizes, self) obj.hide() row = self.rowCount() self.insertRow(row) @@ -450,18 +535,112 @@ def __init__( def values(self) -> dict[str, Any]: result = copy.deepcopy(self.data) for (name, indices), obj in self.objects.items(): + value = obj.value() if indices: - _nested_get(result[name], indices[:-1])[indices[-1]] = obj.value() + _nested_get(result[name], indices[:-1])[indices[-1]] = value + if value: + result[name].assigned.add(tuple(index + 1 for index in indices)) else: - result[name] = obj.value() + result[name] = value return result def reset(self) -> None: - defaults = {name: suggestion(schema, self.sizes) for name, schema in self.schemas.items()} - for (name, indices), obj in self.objects.items(): + for name, data in self.data.items(): + if isinstance(data, InputArray): + shape = array_shape(data) + data.assigned = ( + set(product(*(range(1, n + 1) for n in shape))) + if _seeded(self.schemas[name]) else set() + ) + for obj in self.objects.values(): obj.reset(self.sizes) - for component, value in _nested_get(defaults[name], indices).items(): - obj.rows[component].set_value(value, self.sizes) + + +class ReferencedArrayTable(QTableWidget): + """One row per referenced one-dimensional parameter array.""" + + def __init__( + self, + schemas: Mapping[str, Mapping[str, Any]], + values: Mapping[str, Any], + sizes: Mapping[str, int], + required: set[str], + parent: QWidget, + *, + fit_arrays: bool, + ) -> None: + super().__init__(parent) + self.schemas, self.sizes = schemas, sizes + self.rows: dict[str, list[FieldRow]] = {} + first = next(iter(schemas.values())) + shape = resolve_shape(first, sizes) + if len(shape) != 1: + raise ValueError("referenced parameter arrays must be one-dimensional") + size = shape[0] + self.setRowCount(len(schemas)) + self.setColumnCount(size) + self.setHorizontalHeaderLabels( + axis_labels(first, 1, size) or [str(index) for index in range(1, size + 1)] + ) + cast(QHeaderView, self.horizontalHeader()).setSectionResizeMode( + QHeaderView.ResizeToContents + ) + for row_index, (name, schema) in enumerate(schemas.items()): + item = schema["items"] + raw = values.get(name, MISSING) + saved = raw is not MISSING + dense = initial_array( + schema, + sizes, + raw if saved else suggestion(schema, sizes), + suggestion(item, sizes), + strict=saved and not fit_arrays, + resize=saved and fit_arrays, + ) + positions = {(index,) for index in range(1, size + 1)} + assigned = ( + set(raw.assigned) if isinstance(raw, InputArray) + else positions if saved else set() + ) & positions + if _seeded(schema): + assigned.update(positions) + label = QTableWidgetItem(name + (" *" if name.lower() in required else "")) + label.setToolTip(str(schema.get("description", schema.get("title", name)))) + self.setVerticalHeaderItem(row_index, label) + cells = [] + for index in range(size): + value = dense[index] if (index + 1,) in assigned else MISSING + cell = FieldRow(name, item, value, sizes, self) + self.setCellWidget(row_index, index, cell) + cells.append(cell) + self.rows[name] = cells + self.resizeRowsToContents() + header = cast(QHeaderView, self.horizontalHeader()) + self.setMinimumHeight( + min(400, header.height() + sum(self.rowHeight(i) for i in range(self.rowCount())) + 4) + ) + + def values(self) -> dict[str, InputArray]: + result = {} + for name, cells in self.rows.items(): + assigned = { + (index,) + for index, cell in enumerate(cells, 1) + if cell.value() is not MISSING + } + if assigned: + result[name] = InputArray([cell.field.value() for cell in cells], assigned) + return result + + def reset(self) -> None: + for name, cells in self.rows.items(): + schema = self.schemas[name] + defaults = suggestion(schema, self.sizes) + for index, cell in enumerate(cells): + if _seeded(schema): + cell.set_value(defaults[index], self.sizes) + else: + cell.reset(self.sizes) class NamelistForm(QWidget): @@ -486,14 +665,28 @@ def __init__( required = {item.lower() for item in schema.get("required", []) if isinstance(item, str)} layout = QFormLayout(self) self.rows: dict[str, FieldRow] = {} - self.tables: list[DerivedTable] = [] + self.tables: list[DerivedTable | ReferencedArrayTable] = [] groups: dict[tuple[str, ...], dict[str, Mapping[str, Any]]] = {} identities = {} for name, child in properties.items(): item = child.get("items", {}) if child.get("type") == "array" else child - origin = item.get(DERIVED_REF_ORIGIN_KEY) - if item.get("type") == "object" and origin: - identity = tuple(origin["identity"]) + origin = item.get(DERIVED_REF_ORIGIN_KEY) if item.get("type") == "object" else None + identity = tuple(origin["identity"]) if origin else () + if child.get("type") == "array" and item.get("type") in { + "boolean", "integer", "number", "string" + }: + shape = resolve_shape(child, sizes) + labels = axis_labels(child, 1, shape[0]) if len(shape) == 1 else None + if labels: + identity = ( + "array", + str(shape[0]), + str(item["type"]), + str(item.get("x-fortran-kind")), + str(item.get("x-fortran-len")), + *labels, + ) + if identity: identities[name] = identity groups.setdefault(identity, {})[name] = child for name, child in properties.items(): @@ -503,7 +696,12 @@ def __init__( children = groups.get(identities.get(name, ())) if children: if name == next(iter(children)): - table = DerivedTable( + table_type = ( + ReferencedArrayTable + if child.get("type") == "array" and child["items"].get("type") != "object" + else DerivedTable + ) + table = table_type( children, source, sizes, required, self, fit_arrays=fit_arrays ) layout.addRow(table) diff --git a/src/nml_tools/gui/model.py b/src/nml_tools/gui/model.py index 3ff5c23..89ab933 100644 --- a/src/nml_tools/gui/model.py +++ b/src/nml_tools/gui/model.py @@ -32,6 +32,14 @@ MISSING = object() +class InputArray(list): + """Dense editor values with the indices selected for namelist output.""" + + def __init__(self, values: list[Any], assigned: set[tuple[int, ...]]): + super().__init__(values) + self.assigned = assigned + + def suggestion(schema: Mapping[str, Any], sizes: Mapping[str, int]) -> Any: """Return the deterministic editable value used for an unset schema field.""" examples = schema.get("examples") @@ -352,7 +360,7 @@ def _evaluated_array( target[index] = value elif component in components: target[index][components[component]] = value - return result + return InputArray(result, {coordinates for coordinates, _, _ in states}) def _component_names(schema: Mapping[str, Any]) -> dict[str, str]: @@ -439,6 +447,7 @@ def _normalize_value( items = schema.get("items") if not isinstance(items, Mapping): raise ValueError(f"array '{path}' must define object items") + assigned = value.assigned if isinstance(value, InputArray) else None value = initial_array(schema, sizes, value, suggestion(items, sizes), strict=True) def normalize_items(node: Any, indices: tuple[int, ...] = ()) -> Any: @@ -447,10 +456,13 @@ def normalize_items(node: Any, indices: tuple[int, ...] = ()) -> Any: normalize_items(item, (*indices, index)) for index, item in enumerate(node, start=1) ] + if assigned is not None and indices not in assigned: + return node suffix = "".join(f"[{index}]" for index in indices) return _normalize_value(node, items, sizes, f"{path}{suffix}") - return normalize_items(value) + normalized = normalize_items(value) + return InputArray(normalized, assigned) if assigned is not None else normalized if kind == "object": if not isinstance(value, Mapping): raise ValueError(f"'{path}' must be an object") @@ -672,12 +684,15 @@ def import_profile( def _assignments(name: str, value: Any, schema: Mapping[str, Any]) -> Iterable[str]: if schema["type"] == "array": + assigned = value.assigned if isinstance(value, InputArray) else None def elements(node: Any, indices: tuple[int, ...] = ()) -> Iterable[str]: if isinstance(node, list): for index, child in enumerate(node, 1): yield from elements(child, (*indices, index)) else: + if assigned is not None and indices not in assigned: + return suffix = ",".join(map(str, indices)) yield from _assignments(f"{name}({suffix})", node, schema["items"]) @@ -686,6 +701,8 @@ def elements(node: Any, indices: tuple[int, ...] = ()) -> Iterable[str]: for component, child in value.items(): yield from _assignments(f"{name}%{component}", child, schema["properties"][component]) else: + if schema["type"] == "string" and schema.get("format") == "file-path" and value == "": + return category = "real" if schema["type"] == "number" else schema["type"] yield f" {name} = {_format_scalar_default(value, None, category)}" From b98f0951ccb8d3bcf33ee4e7865e7102be0d08de Mon Sep 17 00:00:00 2001 From: Sanjeev Bashyal Date: Fri, 25 Sep 2026 23:37:28 +0200 Subject: [PATCH 4/6] added PathField and DateTimeField for format-data: file path and date time --- src/nml_tools/gui/app.py | 1 + src/nml_tools/gui/fields.py | 132 ++++++++++++++++++++++++++++++++++++ tests/test_gui_widgets.py | 77 +++++++++++++++++++++ 3 files changed, 210 insertions(+) diff --git a/src/nml_tools/gui/app.py b/src/nml_tools/gui/app.py index ba5fa8f..bad6ab2 100644 --- a/src/nml_tools/gui/app.py +++ b/src/nml_tools/gui/app.py @@ -89,6 +89,7 @@ def __init__( values.get(page.name), sizes, fit_arrays=fit_arrays, + output_root=project.output_root, ) scroll = QScrollArea(self) scroll.setWidgetResizable(True) diff --git a/src/nml_tools/gui/fields.py b/src/nml_tools/gui/fields.py index 5dde006..32e26ef 100644 --- a/src/nml_tools/gui/fields.py +++ b/src/nml_tools/gui/fields.py @@ -4,13 +4,18 @@ import copy import math +import os from collections.abc import Mapping from itertools import product +from pathlib import Path from typing import Any, cast +from qtpy.QtCore import QDateTime from qtpy.QtWidgets import ( QCheckBox, QComboBox, + QDateTimeEdit, + QFileDialog, QFormLayout, QGroupBox, QHBoxLayout, @@ -62,6 +67,85 @@ def accept(self) -> None: return DerivedArrayEditor(parent) +def _output_root(widget: QWidget) -> Path: + current: QWidget | None = widget + while current is not None: + root = getattr(current, "output_root", None) + if isinstance(root, Path): + return root + current = current.parentWidget() + return Path.cwd() + + +def _relative_path(path: str, widget: QWidget) -> str: + return Path(os.path.relpath(path, _output_root(widget))).as_posix() + + +def _parse_date_time(value: Any) -> QDateTime: + text = str(value) + for pattern in ( + "yyyy-MM-dd HH:mm:ss", + "yyyy-MM-dd HH:mm", + "yyyy-MM-dd HH", + "yyyy-MM-dd", + "yyyy-MM-ddTHH:mm:ss", + "yyyy-MM-ddTHH:mm", + ): + result = QDateTime.fromString(text, pattern) + if result.isValid(): + return result + if not text: + return QDateTime.fromString("2000-01-01 00:00", "yyyy-MM-dd HH:mm") + raise ValueError(f"'{text}' is not a valid date-time") + + +def _add_path_array_controls(editor: Any, owner: QWidget) -> tuple[Any, Any, Any]: + line = QLineEdit(editor) + browse = QPushButton("...", editor) + update = QPushButton("Update selected", editor) + controls = QHBoxLayout() + controls.addWidget(line, 1) + controls.addWidget(browse) + controls.addWidget(update) + editor.arraywidget.layout().addLayout(controls) + + def choose() -> None: + path, _ = QFileDialog.getOpenFileName(editor, "Select file", str(_output_root(owner))) + if path: + line.setText(_relative_path(path, owner)) + + def apply() -> None: + model = editor.arraywidget.model + for index in editor.arraywidget.view.selectedIndexes(): + model.setData(index, line.text()) + + browse.clicked.connect(choose) + update.clicked.connect(apply) + return line, browse, update + + +def _install_date_time_delegate(editor: Any) -> None: + from guidata.widgets.arrayeditor.editorwidget import ( # type: ignore[import-untyped] + ArrayDelegate, + ) + + class DateTimeDelegate(ArrayDelegate): # type: ignore[misc] + def createEditor(self, parent: QWidget, option: Any, index: Any) -> QDateTimeEdit: + control = QDateTimeEdit(parent) + control.setCalendarPopup(True) + control.setDisplayFormat("yyyy-MM-dd HH:mm") + return control + + def setEditorData(self, control: QDateTimeEdit, index: Any) -> None: + control.setDateTime(_parse_date_time(index.model().data(index))) + + def setModelData(self, control: QDateTimeEdit, model: Any, index: Any) -> None: + model.setData(index, control.dateTime().toString("yyyy-MM-dd HH:mm")) + + view = editor.arraywidget.view + view.setItemDelegate(DateTimeDelegate(view.model().get_array().dtype, view)) + + def _seeded(schema: Mapping[str, Any]) -> bool: if "default" in schema or schema.get("examples"): return True @@ -136,6 +220,44 @@ def reset(self, sizes: Mapping[str, int]) -> None: self.modified = False +class PathField(ScalarField): + def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None = None): + super().__init__(schema, value, parent) + self.browse = QPushButton("...", self) + self.layout().addWidget(self.browse) + self.browse.clicked.connect(self._browse) + + def _browse(self) -> None: + path, _ = QFileDialog.getOpenFileName(self, "Select file", str(_output_root(self))) + if path: + self.set_value(_relative_path(path, self)) + + +class DateTimeField(ScalarField): + def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None = None): + QWidget.__init__(self, parent) + self.schema = schema + layout = QHBoxLayout(self) + layout.setContentsMargins(0, 0, 0, 0) + self.control = QDateTimeEdit(self) + self.control.setCalendarPopup(True) + self.control.setDisplayFormat("yyyy-MM-dd HH:mm") + layout.addWidget(self.control) + self.set_value(value) + self.modified = False + self.control.dateTimeChanged.connect(lambda *_: setattr(self, "modified", True)) + + def set_value(self, value: Any) -> None: + self._original = str(value) + self.control.setDateTime(_parse_date_time(value)) + self.modified = True + + def value(self) -> str: + if not self.modified: + return self._original + return self.control.dateTime().toString("yyyy-MM-dd HH:mm") + + class ObjectField(QGroupBox): def __init__( self, @@ -299,6 +421,10 @@ def _edit(self) -> None: variable_size=deferred, ): return + if self.items.get("format") == "file-path" and not derived: + _add_path_array_controls(editor, self) + if self.items.get("format") == "date-time" and not derived: + _install_date_time_delegate(editor) if _exec(editor) != _accepted(editor): return edited = editor.get_value() @@ -654,10 +780,12 @@ def __init__( parent: QWidget | None = None, *, fit_arrays: bool = False, + output_root: Path | None = None, ): super().__init__(parent) self.schema = schema self.sizes = sizes + self.output_root = (output_root or Path.cwd()).resolve() properties = schema.get("properties") if not isinstance(properties, Mapping): raise ValueError("namelist schema must define object 'properties'") @@ -749,6 +877,10 @@ def _field_widget( return ArrayField(name, schema, value, sizes, parent, fit_existing=fit_arrays) if kind == "object": return ObjectField(schema, value, sizes, parent, fit_arrays=fit_arrays) + if kind == "string" and schema.get("format") == "file-path": + return PathField(schema, value, parent) + if kind == "string" and schema.get("format") == "date-time": + return DateTimeField(schema, value, parent) return ScalarField(schema, value, parent) diff --git a/tests/test_gui_widgets.py b/tests/test_gui_widgets.py index d67e96b..3c16261 100644 --- a/tests/test_gui_widgets.py +++ b/tests/test_gui_widgets.py @@ -85,6 +85,50 @@ def test_derived_singletons_use_inline_object_fields(application, project): assert scalar.value() == {"year": 2021} +def test_path_and_date_time_fields(application, tmp_path, monkeypatch): + from qtpy.QtCore import QDateTime + + from nml_tools.gui import fields + from nml_tools.gui.fields import DateTimeField, PathField + + schema = { + "type": "object", + "properties": { + "path": {"type": "string", "format": "file-path"}, + "paths": { + "type": "array", + "x-fortran-shape": "n", + "items": {"type": "string", "format": "file-path"}, + }, + "when": {"type": "string", "format": "date-time"}, + }, + } + form = NamelistForm( + schema, + {"path": "old.nc", "paths": ["one.nc"], "when": "2025-01-01"}, + {"n": 1}, + output_root=tmp_path, + ) + path = form.rows["path"].field + assert isinstance(path, PathField) + assert isinstance(form.rows["paths"].field.inline, PathField) + monkeypatch.setattr( + fields.QFileDialog, + "getOpenFileName", + lambda *args: (str(tmp_path / "data" / "input.nc"), ""), + ) + path.browse.click() + assert path.value() == "data/input.nc" + + date_time = form.rows["when"].field + assert isinstance(date_time, DateTimeField) + assert date_time.value() == "2025-01-01" + date_time.control.setDateTime( + QDateTime.fromString("2026-02-03 04:05", "yyyy-MM-dd HH:mm") + ) + assert date_time.value() == "2026-02-03 04:05" + + def test_dialog_load_overlay_save_reload_and_dimension_changes(application, project, monkeypatch): errors = [] monkeypatch.setattr(QMessageBox, "critical", lambda *args: errors.append(args[-1])) @@ -159,6 +203,39 @@ def test_guidata_derived_edits_commit(application): editor.close() +def test_guidata_path_update_and_date_time_editor(application): + np = pytest.importorskip("numpy") + pytest.importorskip("guidata") + from guidata.widgets.arrayeditor import ArrayEditor + from qtpy.QtCore import QDateTime + from qtpy.QtWidgets import QDateTimeEdit + + from nml_tools.gui.fields import _add_path_array_controls, _install_date_time_delegate + + data = np.array([["a.nc", "b.nc"]], dtype="U1024") + editor = ArrayEditor(None) + try: + assert editor.setup_and_check(data) + line, _, update = _add_path_array_controls(editor, editor) + editor.arraywidget.view.selectAll() + line.setText("data/input.nc") + update.click() + model = editor.arraywidget.model + assert model.get_value((0, 0)) == model.get_value((0, 1)) == "data/input.nc" + + _install_date_time_delegate(editor) + view = editor.arraywidget.view + index = view.model().index(0, 0) + delegate = view.itemDelegate() + control = delegate.createEditor(view, None, index) + assert isinstance(control, QDateTimeEdit) + control.setDateTime(QDateTime.fromString("2026-02-03 04:05", "yyyy-MM-dd HH:mm")) + delegate.setModelData(control, view.model(), index) + assert model.get_value((0, 0)) == "2026-02-03 04:05" + finally: + editor.close() + + def test_imported_profile_can_choose_its_output_and_keep_loaded_values( application, project, monkeypatch ): From 26d4a34360d7c15be07a868e2191c77013ab4540 Mon Sep 17 00:00:00 2001 From: Sanjeev Bashyal Date: Sat, 26 Sep 2026 09:25:56 +0200 Subject: [PATCH 5/6] Fix: for multi file profiles, it should be the first file profile to be selected at the end after loading all --- src/nml_tools/gui/app.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/nml_tools/gui/app.py b/src/nml_tools/gui/app.py index bad6ab2..1e3f65f 100644 --- a/src/nml_tools/gui/app.py +++ b/src/nml_tools/gui/app.py @@ -511,6 +511,8 @@ def _run_configuration(self, config: ProfileConfigTab) -> None: prepared.append(editor) for editor in prepared: self._put_editor(editor) + if prepared: + self.tabs.setCurrentWidget(prepared[0]) self.dimensions = dimensions if config.builder or config.source_path is not None: self._remove_tab(config) From f8d46cb148729b480642af2cae526844f7317b9c Mon Sep 17 00:00:00 2001 From: Sanjeev Bashyal Date: Mon, 5 Oct 2026 14:26:48 +0200 Subject: [PATCH 6/6] Improvement: added Spinbox for type integer instead of lineEdit --- src/nml_tools/gui/fields.py | 19 ++++++++++++++++++- tests/test_gui_widgets.py | 22 ++++++++++++++++++++++ 2 files changed, 40 insertions(+), 1 deletion(-) diff --git a/src/nml_tools/gui/fields.py b/src/nml_tools/gui/fields.py index 32e26ef..c17e78a 100644 --- a/src/nml_tools/gui/fields.py +++ b/src/nml_tools/gui/fields.py @@ -24,6 +24,7 @@ QLineEdit, QMessageBox, QPushButton, + QSpinBox, QTableWidget, QTableWidgetItem, QWidget, @@ -165,7 +166,7 @@ def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None layout.setContentsMargins(0, 0, 0, 0) enum = schema.get("enum") kind = schema.get("type") - control: QComboBox | QCheckBox | QLineEdit + control: QComboBox | QCheckBox | QSpinBox | QLineEdit if isinstance(enum, list) and enum: combo = QComboBox(self) for item in enum: @@ -173,6 +174,17 @@ def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None control = combo elif kind == "boolean": control = QCheckBox(self) + elif kind == "integer": + lower = int(schema.get("minimum", schema.get("exclusiveMinimum", -2_147_483_648))) + upper = int(schema.get("maximum", schema.get("exclusiveMaximum", 2_147_483_647))) + lower += "exclusiveMinimum" in schema + upper -= "exclusiveMaximum" in schema + if -2_147_483_648 <= lower <= int(value) <= upper <= 2_147_483_647: + spin = QSpinBox(self) + spin.setRange(lower, upper) + control = spin + else: + control = QLineEdit(self) else: control = QLineEdit(self) self.control = control @@ -182,6 +194,7 @@ def __init__(self, schema: Mapping[str, Any], value: Any, parent: QWidget | None signal = ( control.textEdited if isinstance(control, QLineEdit) else control.toggled if isinstance(control, QCheckBox) + else control.valueChanged if isinstance(control, QSpinBox) else control.currentIndexChanged ) signal.connect(lambda *_: setattr(self, "modified", True)) @@ -192,6 +205,8 @@ def set_value(self, value: Any) -> None: self.control.setCurrentIndex(max(index, 0)) elif isinstance(self.control, QCheckBox): self.control.setChecked(bool(value)) + elif isinstance(self.control, QSpinBox): + self.control.setValue(value) else: self.control.setText(str(value)) self.modified = True @@ -201,6 +216,8 @@ def value(self) -> Any: return self.control.currentData() if isinstance(self.control, QCheckBox): return self.control.isChecked() + if isinstance(self.control, QSpinBox): + return self.control.value() text = self.control.text() kind = self.schema.get("type") try: diff --git a/tests/test_gui_widgets.py b/tests/test_gui_widgets.py index 3c16261..3761e7e 100644 --- a/tests/test_gui_widgets.py +++ b/tests/test_gui_widgets.py @@ -73,6 +73,28 @@ def test_singletons_keep_array_values_and_restore_schema_defaults(application): assert not deferred.field.button.isHidden() +def test_numeric_fields_use_appropriate_controls(application): + from qtpy.QtWidgets import QLineEdit, QSpinBox + + integer = ScalarField({"type": "integer", "minimum": -3, "maximum": 9}, 4) + assert isinstance(integer.control, QSpinBox) + assert (integer.control.minimum(), integer.control.maximum(), integer.value()) == (-3, 9, 4) + singleton = FieldRow( + "values", {"type": "array", "x-fortran-shape": 1, "items": {"type": "integer"}}, [2], {} + ) + assert isinstance(singleton.field.inline.control, QSpinBox) + + number = ScalarField({"type": "number", "minimum": 0.0, "maximum": 1.0}, 0.25) + assert isinstance(number.control, QLineEdit) + assert number.value() == 0.25 + number.control.setText("not-a-number") + with pytest.raises(ValueError, match="not a valid number"): + number.value() + + wide = ScalarField({"type": "integer", "maximum": 2**40}, 4) + assert isinstance(wide.control, QLineEdit) + + def test_derived_singletons_use_inline_object_fields(application, project): schema = project.namelists[0].schema["properties"]["periods"] row = FieldRow("periods", schema, [{"year": 2020}], {"n": 1})