diff --git a/.github/workflows/testing.yaml b/.github/workflows/testing.yaml index 8c208926c..ae4c79183 100644 --- a/.github/workflows/testing.yaml +++ b/.github/workflows/testing.yaml @@ -22,7 +22,6 @@ jobs: steps: - uses: actions/checkout@v4 - # general Python setup - name: Set up Python ${{ matrix.py-version }} uses: actions/setup-python@v5 with: @@ -37,11 +36,9 @@ jobs: python -VV python -m pip install --upgrade pip setuptools wheel - # install testing - name: Install package and test deps run: | - pip install .[testing] # install the package and the testing deps + pip install .[testing] - name: Test with pytest - run: | - pytest -s --no-header --no-summary -q --no-cov + run: pytest -q --no-cov -o log_cli=false diff --git a/activity_browser/app/main.py b/activity_browser/app/main.py index 8c50c5aa2..c2b86e903 100644 --- a/activity_browser/app/main.py +++ b/activity_browser/app/main.py @@ -115,7 +115,7 @@ def apply_settings(self, load=False): else: hint = QtCore.Qt.ColorScheme.Unknown - app.application.styleHints().setColorScheme(hint) + app.application.apply_color_scheme(hint) # apply pane tab position position = app.settings["appearance"]["pane_tab_position"] diff --git a/activity_browser/app/menu_bar.py b/activity_browser/app/menu_bar.py index ebb57c53b..2e64cc80d 100644 --- a/activity_browser/app/menu_bar.py +++ b/activity_browser/app/menu_bar.py @@ -7,7 +7,7 @@ from qtpy.QtCore import QSize, QUrl, Qt from activity_browser import app -from activity_browser.bwutils.commontasks import get_templates +from activity_browser.bwutils.commontasks import fetch_remote_projects, get_templates from ..ui.icons import qicons @@ -111,8 +111,7 @@ def __init__(self, parent=None): def get_projects(self): if not self.remote_projects: - from bw2io.remote import get_projects - ProjectNewTemplateMenu.remote_projects = get_projects() + ProjectNewTemplateMenu.remote_projects = fetch_remote_projects() return self.remote_projects diff --git a/activity_browser/app/pages/lca_results/LCA_results.py b/activity_browser/app/pages/lca_results/LCA_results.py index 64af55f27..2b2411305 100644 --- a/activity_browser/app/pages/lca_results/LCA_results.py +++ b/activity_browser/app/pages/lca_results/LCA_results.py @@ -252,6 +252,7 @@ def __init__(self, parent=None): self.layout = QtWidgets.QVBoxLayout() self.setLayout(self.layout) configure_lca_tab_layout(self.layout) + app.application.theme_changed.connect(self.update_tab) def add_tab_header( self, title: str, help_tooltip: Optional[str] = None @@ -277,7 +278,6 @@ def build_tab_body( QtWidgets.QSizePolicy.Policy.Expanding, ) widget.setMinimumWidth(0) - widget.setStyleSheet("background-color: white;") self.pt_layout.setContentsMargins(0, 0, 0, 0) self.pt_layout.setAlignment(alignment) widget.setLayout(self.pt_layout) diff --git a/activity_browser/app/pages/lca_results/plots.py b/activity_browser/app/pages/lca_results/plots.py index 44b3eddee..bdfc39f75 100644 --- a/activity_browser/app/pages/lca_results/plots.py +++ b/activity_browser/app/pages/lca_results/plots.py @@ -11,6 +11,7 @@ from matplotlib.gridspec import GridSpec from matplotlib.patches import Patch +import matplotlib.pyplot as plt import numpy as np import pandas as pd from bw2data import methods @@ -318,13 +319,14 @@ def _prepare_data(self, df: pd.DataFrame, unit: str | None) -> _ContributionFram @staticmethod def _draw_net_markers(ax, dfp: pd.DataFrame, *, horizontal: bool) -> None: - """Black dot at bar tip when positive and negative stacks do not cancel.""" + """Dot at bar tip when positive and negative stacks do not cancel.""" + edge_color = plt.rcParams["axes.edgecolor"] kw = dict( markersize=5, marker="o", linestyle="none", - markerfacecolor="black", - markeredgecolor="black", + markerfacecolor=edge_color, + markeredgecolor=edge_color, ) for i in range(dfp.shape[1]): s = dfp.iloc[:, i] @@ -480,7 +482,7 @@ def _render(self) -> None: Patch(color=self.gsa_type_color(t), label=t) for t in dfp[GSA_TYPE_COLUMN].drop_duplicates() ] - err_kw = dict(capsize=3, ecolor="#333333") + err_kw = dict(capsize=3, ecolor=plt.rcParams["axes.edgecolor"]) legend_ratio = 0.18 if handles else 0.0 value_label = "Delta sensitivity index" diff --git a/activity_browser/app/pages/settings/appearance.py b/activity_browser/app/pages/settings/appearance.py index 9340be7a4..3d6fa8c1f 100644 --- a/activity_browser/app/pages/settings/appearance.py +++ b/activity_browser/app/pages/settings/appearance.py @@ -1,9 +1,19 @@ # -*- coding: utf-8 -*- -from loguru import logger from qtpy import QtWidgets from activity_browser.app import settings from activity_browser.app.pages.settings.base import BaseSettingsChapter +from activity_browser.ui.widgets.plot import DEFAULT_PLOT_PALETTE, PLOT_PALETTES + + +def _labeled_group(title: str, rows: list[tuple[str, QtWidgets.QWidget]]) -> QtWidgets.QGroupBox: + group = QtWidgets.QGroupBox(title) + grid = QtWidgets.QGridLayout() + for i, (label, widget) in enumerate(rows): + grid.addWidget(QtWidgets.QLabel(label), i, 0) + grid.addWidget(widget, i, 1) + group.setLayout(grid) + return group class AppearanceSettingsChapter(BaseSettingsChapter): @@ -14,106 +24,95 @@ class AppearanceSettingsChapter(BaseSettingsChapter): "light": "Light theme", "dark": "Dark theme", } - pane_tab_position_map = { "top": "Top", "bottom": "Bottom", "left": "Left", "right": "Right", } - + def __init__(self, parent=None): super().__init__(parent) - - # Theme selector self.theme_combo = QtWidgets.QComboBox() - - # Pane tab position selector + self.palette_combo = QtWidgets.QComboBox() self.pane_tab_position_combo = QtWidgets.QComboBox() - - # Database products as cards checkbox self.database_products_as_cards = QtWidgets.QCheckBox("Show database contents as cards") self.database_products_as_cards.setToolTip( "When enabled, the database process list uses a card layout instead of a detailed table." ) - self.build_layout() self.connect_signals() self.reset() - + def connect_signals(self): - """Connect signals and slots.""" - # Emit changed signal when settings change - self.theme_combo.currentTextChanged.connect(lambda: self.changed.emit()) - self.pane_tab_position_combo.currentTextChanged.connect(lambda: self.changed.emit()) + for widget in ( + self.theme_combo, + self.palette_combo, + self.pane_tab_position_combo, + ): + widget.currentTextChanged.connect(lambda: self.changed.emit()) self.database_products_as_cards.checkStateChanged.connect(lambda _: self.changed.emit()) - + def build_layout(self): - """Build the chapter layout.""" layout = QtWidgets.QVBoxLayout() - - # Theme section - theme_group = QtWidgets.QGroupBox("Theme") - theme_layout = QtWidgets.QGridLayout() - theme_layout.addWidget(QtWidgets.QLabel("Theme:"), 0, 0) - theme_layout.addWidget(self.theme_combo, 0, 1) - theme_group.setLayout(theme_layout) - - # Pane tab position section - pane_tab_group = QtWidgets.QGroupBox("Pane Tab Position") - pane_tab_layout = QtWidgets.QGridLayout() - pane_tab_layout.addWidget(QtWidgets.QLabel("Position:"), 0, 0) - pane_tab_layout.addWidget(self.pane_tab_position_combo, 0, 1) - pane_tab_group.setLayout(pane_tab_layout) - + layout.addWidget(_labeled_group("Theme", [("Theme:", self.theme_combo)])) + layout.addWidget(_labeled_group("Plots", [("Palette:", self.palette_combo)])) + layout.addWidget( + _labeled_group("Pane Tab Position", [("Position:", self.pane_tab_position_combo)]) + ) database_group = QtWidgets.QGroupBox("Database pane") - database_layout = QtWidgets.QVBoxLayout() - database_layout.addWidget(self.database_products_as_cards) - database_group.setLayout(database_layout) - - layout.addWidget(theme_group) - layout.addWidget(pane_tab_group) + db_layout = QtWidgets.QVBoxLayout() + db_layout.addWidget(self.database_products_as_cards) + database_group.setLayout(db_layout) layout.addWidget(database_group) layout.addStretch() - self.setLayout(layout) - - # --- Settings management methods --- # + + @staticmethod + def _saved_palette() -> str: + name = settings["appearance"].get("plot_palette", DEFAULT_PLOT_PALETTE) + return name if name in PLOT_PALETTES else DEFAULT_PLOT_PALETTE + def reset(self): - """(Re)set to initial values.""" + appearance = settings["appearance"] self.theme_combo.clear() self.theme_combo.addItems(self.theme_map.values()) - self.theme_combo.setCurrentText(self.theme_map.get(settings["appearance"]["theme"], "System default")) - + self.theme_combo.setCurrentText(self.theme_map.get(appearance["theme"], "System default")) + self.palette_combo.clear() + self.palette_combo.addItems(PLOT_PALETTES) + self.palette_combo.setCurrentText(self._saved_palette()) self.pane_tab_position_combo.clear() self.pane_tab_position_combo.addItems(self.pane_tab_position_map.values()) - self.pane_tab_position_combo.setCurrentText(self.pane_tab_position_map.get(settings["appearance"]["pane_tab_position"], "Bottom")) - - self.database_products_as_cards.setChecked( - bool(settings["appearance"].get("database_products_as_cards", False)) + self.pane_tab_position_combo.setCurrentText( + self.pane_tab_position_map.get(appearance["pane_tab_position"], "Bottom") ) + self.database_products_as_cards.setChecked(bool(appearance.get("database_products_as_cards", False))) def has_changes(self): - """Check if there are unsaved changes.""" - current_state = { - 'theme': self.theme_combo.currentText(), - 'pane_tab_position': self.pane_tab_position_combo.currentText(), - 'database_products_as_cards': self.database_products_as_cards.isChecked(), + appearance = settings["appearance"] + return { + "theme": self.theme_combo.currentText(), + "plot_palette": self.palette_combo.currentText(), + "pane_tab_position": self.pane_tab_position_combo.currentText(), + "database_products_as_cards": self.database_products_as_cards.isChecked(), + } != { + "theme": self.theme_map.get(appearance["theme"], "System default"), + "plot_palette": self._saved_palette(), + "pane_tab_position": self.pane_tab_position_map.get( + appearance["pane_tab_position"], "Bottom" + ), + "database_products_as_cards": bool(appearance.get("database_products_as_cards", False)), } - initial_state = { - 'theme': self.theme_map.get(settings["appearance"]["theme"], "System default"), - 'pane_tab_position': self.pane_tab_position_map.get(settings["appearance"]["pane_tab_position"], "Bottom"), - 'database_products_as_cards': bool(settings["appearance"].get("database_products_as_cards", False)), - } - return current_state != initial_state - - def set_settings(self): - """Save appearance settings.""" - new_theme = self.theme_combo.currentText() - settings["appearance"]["theme"] = [key for key, value in self.theme_map.items() if value == new_theme][0] - - new_pane_position = self.pane_tab_position_combo.currentText() - settings["appearance"]["pane_tab_position"] = [key for key, value in self.pane_tab_position_map.items() if value == new_pane_position][0] - - settings["appearance"]["database_products_as_cards"] = self.database_products_as_cards.isChecked() + def set_settings(self): + appearance = settings["appearance"] + appearance["theme"] = next( + k for k, v in self.theme_map.items() if v == self.theme_combo.currentText() + ) + appearance["plot_palette"] = self.palette_combo.currentText() + appearance["pane_tab_position"] = next( + k + for k, v in self.pane_tab_position_map.items() + if v == self.pane_tab_position_combo.currentText() + ) + appearance["database_products_as_cards"] = self.database_products_as_cards.isChecked() diff --git a/activity_browser/app/pages/settings/project_manager.py b/activity_browser/app/pages/settings/project_manager.py index 925550825..684cf2f88 100644 --- a/activity_browser/app/pages/settings/project_manager.py +++ b/activity_browser/app/pages/settings/project_manager.py @@ -4,10 +4,9 @@ from qtpy import QtWidgets, QtGui import bw2data as bd -from bw2io import remote from activity_browser import app, ui -from activity_browser.bwutils.commontasks import get_templates +from activity_browser.bwutils.commontasks import fetch_remote_projects, get_templates from activity_browser.ui import widgets, core from .base import BaseSettingsChapter @@ -94,7 +93,7 @@ def build_template_df(self) -> pd.DataFrame: data = [] templates = get_templates() - remote_templates = remote.get_projects() + remote_templates = fetch_remote_projects() for name in sorted(templates): data.append({ diff --git a/activity_browser/app/pages/welcome.py b/activity_browser/app/pages/welcome.py index d2c8ea9d4..4fd0c1edf 100644 --- a/activity_browser/app/pages/welcome.py +++ b/activity_browser/app/pages/welcome.py @@ -35,7 +35,24 @@ def __init__(self, parent=None): self.setLayout(self.vl) self.bridge.ready.connect(self.update_welcome) - app.signals.project.changed.connect(lambda: self.page.load(self.url)) + app.application.theme_changed.connect(self._reload_page) + app.signals.project.changed.connect(self._reload_page) + self.page.loadFinished.connect(self._on_load_finished) + + def _reload_page(self, *_args) -> None: + self.page.load(self.url) + + def _on_load_finished(self, ok: bool) -> None: + if not ok: + return + scheme = ( + "dark" + if app.application.styleHints().colorScheme() == QtCore.Qt.ColorScheme.Dark + else "light" + ) + self.page.runJavaScript( + f'document.documentElement.style.colorScheme = "{scheme}";' + ) def update_welcome(self): projects = projects_by_last_opened() diff --git a/activity_browser/bwutils/commontasks.py b/activity_browser/bwutils/commontasks.py index 32b568324..8aa9be077 100644 --- a/activity_browser/bwutils/commontasks.py +++ b/activity_browser/bwutils/commontasks.py @@ -629,6 +629,20 @@ def get_templates() -> dict: return collection + +def fetch_remote_projects() -> dict: + """Remote template catalogue from ``bw2io``; empty dict if unreachable.""" + try: + from bw2io.remote import get_projects + + return get_projects() or {} + except Exception as exc: + from loguru import logger + + logger.warning(f"Could not fetch remote project templates: {exc}") + return {} + + def nodes_to_excel(nodes: list[tuple | int | bd.Node]) -> str: """Convert a list of nodes to an HTML table suitable for Excel.""" from .exporters import ABCSVFormatter diff --git a/activity_browser/bwutils/settings.py b/activity_browser/bwutils/settings.py index a2b19eb82..a609f1b38 100644 --- a/activity_browser/bwutils/settings.py +++ b/activity_browser/bwutils/settings.py @@ -18,6 +18,7 @@ "theme": "default", "pane_tab_position": "bottom", "database_products_as_cards": False, + "plot_palette": "tab20", }, "metadatastore": { "caching_enabled": True, diff --git a/activity_browser/bwutils/uncertainty.py b/activity_browser/bwutils/uncertainty.py index 5f5f22310..9af92b245 100644 --- a/activity_browser/bwutils/uncertainty.py +++ b/activity_browser/bwutils/uncertainty.py @@ -24,8 +24,7 @@ import stats_arrays as sa from bw2data.parameters import ParameterBase from bw2data.proxies import ExchangeProxyBase -from stats_arrays import UncertaintyBase, UndefinedUncertainty -from stats_arrays import uncertainty_choices as uc +from stats_arrays import UncertaintyBase, UndefinedUncertainty, uncertainty_choices as uc # Cleared uncertainty state for remove-uncertainty actions and dialog defaults. EMPTY_UNCERTAINTY = { @@ -38,6 +37,127 @@ "negative": False, } +# Fields that may be left empty; ``stats_arrays`` supplies defaults or ignores them. +OPTIONAL_UNCERTAINTY_FIELDS = { + sa.BetaUncertainty.id: frozenset({"minimum", "maximum"}), + sa.BetaPERTUncertainty.id: frozenset({"scale"}), + sa.DiscreteUniform.id: frozenset({"minimum"}), + sa.StudentsTUncertainty.id: frozenset({"loc", "scale"}), + sa.GeneralizedExtremeValueUncertainty.id: frozenset({"loc", "scale"}), +} + +# Distributions whose arithmetic mean is shown read-only (from ``statistics()``). +DISTRIBUTIONS_WITH_CALCULATED_MEAN = frozenset({ + sa.TriangularUncertainty.id, + sa.UniformUncertainty.id, + sa.DiscreteUniform.id, + sa.BetaUncertainty.id, + sa.BetaPERTUncertainty.id, +}) + +_UNCERTAINTY_VALIDATION_N = 24 +UNCERTAINTY_VALIDATION_N = _UNCERTAINTY_VALIDATION_N + + +def as_scalar(value) -> float: + """One float from a stats_arrays statistic (scalar, 0-d, or 1-element array).""" + if value is None: + return float("nan") + arr = np.asarray(value, dtype=float).ravel() + if arr.size == 0: + return float("nan") + return float(arr[0]) + + +def discrete_uniform_expected(array: np.ndarray) -> float: + """Expected value on integers ``[minimum, maximum)`` (:class:`stats_arrays.DiscreteUniform`).""" + if array is None or len(array) == 0: + return float("nan") + row = array[0] + lo = int(round(as_scalar(row["minimum"]))) + hi = int(round(as_scalar(row["maximum"]))) + if hi < lo: + lo, hi = hi, lo + if hi <= lo: + return float("nan") + return (float(lo) + float(hi) - 1.0) / 2.0 + + +def prepare_uncertainty_dict(info: dict, dist=None) -> dict: + """Merge defaults and apply ``stats_arrays``-specific fixes before validate/sample.""" + data = {**EMPTY_UNCERTAINTY, **(info or {})} + ut_id = int(data.get("uncertainty type", 0) or 0) + if dist is not None: + ut_id = dist.id + data["uncertainty type"] = ut_id + + # Only ξ = 0 (Gumbel) is implemented; ``nan != 0`` would fail validation otherwise. + if ut_id == sa.GeneralizedExtremeValueUncertainty.id: + data["shape"] = 0.0 + if np.isnan(data.get("loc", np.nan)): + data["loc"] = 0.0 + if np.isnan(data.get("scale", np.nan)): + data["scale"] = 1.0 + return data + + +def validate_uncertainty_dict( + info: dict, dist=None, n: int = _UNCERTAINTY_VALIDATION_N +) -> tuple[np.ndarray | None, str | None]: + """Build via ``UncertaintyBase.from_dicts``, then ``validate`` and sample.""" + data = prepare_uncertainty_dict(info, dist) + cls = uc[int(data["uncertainty type"])] + try: + array = UncertaintyBase.from_dicts(data) + cls.validate(array) + cls.random_variables(array, n) + return array, None + except Exception as exc: + msg = str(exc).strip() or exc.__class__.__name__ + return None, msg + + +def uncertainty_dict_is_sampleable(data: dict) -> bool: + """True when ``stats_arrays`` accepts *data* (undefined / no-uncertainty always pass).""" + try: + ut_id = int(data.get("uncertainty type", 0)) + except (TypeError, ValueError): + return False + if ut_id < 0 or ut_id >= len(uc): + return False + dist = uc[ut_id] + if dist.id in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id): + return True + array, _ = validate_uncertainty_dict(data) + return array is not None + + +def uncertainty_statistics_scalar(dist, array: np.ndarray, key: str = "mean") -> float: + """One distribution statistic; works around NumPy 2.x bugs in some ``statistics()`` methods.""" + if array is None or len(array) == 0 or dist is None: + return float("nan") + try: + stats = dist.statistics(array) + val = stats.get(key) + if val in (None, "Not Implemented"): + val = stats.get("mean") + return as_scalar(val) + except (TypeError, ValueError, IndexError): + pass + if dist.id == sa.LognormalUncertainty.id and key == "median": + row = array[0] + sign = -1.0 if bool(row["negative"]) else 1.0 + return sign * float(np.exp(as_scalar(row["loc"]))) + if dist.id == sa.DiscreteUniform.id: + return discrete_uniform_expected(array) + return float("nan") + + +def uncertainty_reference_value(dist, array: np.ndarray) -> float: + """Preview plot reference line (median for lognormal, mean otherwise).""" + stat = "median" if dist.id == sa.LognormalUncertainty.id else "mean" + return uncertainty_statistics_scalar(dist, array, stat) + def uncertainty_type_id(source) -> int: """Return stats_arrays uncertainty type id from a dict, exchange proxy, or int.""" @@ -79,6 +199,8 @@ def standard_uncertainty_fields(ut_id: int) -> list[str]: return ["loc", "scale", "shape"] if ut_id == sa.BetaUncertainty.id: return ["loc", "shape", "minimum", "maximum"] + if ut_id == sa.BetaPERTUncertainty.id: + return ["minimum", "loc", "maximum", "scale"] if ut_id == sa.GeneralizedExtremeValueUncertainty.id: return ["loc", "scale"] return [] @@ -94,11 +216,13 @@ def uncertainty_field_name(ut_id: int, field_key: str) -> str: sa.LognormalUncertainty.id: "Loc (ln(mean))", sa.TriangularUncertainty.id: "Mode", sa.BetaUncertainty.id: "Alpha (α)", + sa.BetaPERTUncertainty.id: "Mean (B)", }.get(ut_id, "Loc / offset" if ut_id in (sa.GammaUncertainty.id, sa.WeibullUncertainty.id) else "Mean / location") if field_key == "scale": return { sa.GammaUncertainty.id: "Scale (θ)", sa.WeibullUncertainty.id: "Scale (λ)", + sa.BetaPERTUncertainty.id: "Lambda (λ)", }.get( ut_id, "Scale (σ)" @@ -120,9 +244,13 @@ def uncertainty_field_name(ut_id: int, field_key: str) -> str: "Shape (k)" if ut_id in (sa.GammaUncertainty.id, sa.WeibullUncertainty.id) else "Shape", ) if field_key == "minimum": - return "Minimum (inclusive)" if ut_id == sa.DiscreteUniform.id else "Minimum" + return "Minimum (A)" if ut_id == sa.BetaPERTUncertainty.id else ( + "Minimum (inclusive)" if ut_id == sa.DiscreteUniform.id else "Minimum" + ) if field_key == "maximum": - return "Maximum (exclusive)" if ut_id == sa.DiscreteUniform.id else "Maximum" + return "Maximum (C)" if ut_id == sa.BetaPERTUncertainty.id else ( + "Maximum (exclusive)" if ut_id == sa.DiscreteUniform.id else "Maximum" + ) return field_key diff --git a/activity_browser/static/startscreen/welcome.html b/activity_browser/static/startscreen/welcome.html index 5d9661203..8a20e1b14 100644 --- a/activity_browser/static/startscreen/welcome.html +++ b/activity_browser/static/startscreen/welcome.html @@ -4,11 +4,16 @@ Welcome - diff --git a/activity_browser/ui/core/application.py b/activity_browser/ui/core/application.py index 674bdf050..b376a8922 100644 --- a/activity_browser/ui/core/application.py +++ b/activity_browser/ui/core/application.py @@ -3,15 +3,32 @@ from loguru import logger from qtpy import QtGui, QtWidgets, QtCore, PYSIDE6 -from qtpy.QtCore import Qt +from qtpy.QtCore import Qt, Signal from qtpy.QtGui import QFontDatabase from activity_browser.static import fonts, icons +_OFFSCREEN_FLAGS = ("--disable-gpu", "--disable-gpu-compositing", "--no-sandbox") + + +def _webengine_flags(*add: str, drop: tuple[str, ...] = ()) -> None: + """Set QTWEBENGINE_CHROMIUM_FLAGS before QtWebEngineQuick.initialize().""" + skip = set(drop) + order = ( + *(_OFFSCREEN_FLAGS if os.environ.get("QT_QPA_PLATFORM") == "offscreen" else ()), + *os.environ.get("QTWEBENGINE_CHROMIUM_FLAGS", "").split(), + *add, + ) + os.environ["QTWEBENGINE_CHROMIUM_FLAGS"] = " ".join( + dict.fromkeys(f for f in order if f and f not in skip) + ) + + class ABApplication(QtWidgets.QApplication): _main_window = None _instance = None + theme_changed = Signal() windows = [] @@ -51,6 +68,8 @@ def set_icon(self): def pyside6_setup(self): from qtpy.QtWebEngineQuick import QtWebEngineQuick + + _webengine_flags() QtWebEngineQuick.initialize() style = QtWidgets.QStyleFactory().create("fusion") @@ -68,14 +87,25 @@ def check_palette(self, color_scheme): plt.style.use("dark_background") - os.environ["QTWEBENGINE_CHROMIUM_FLAGS"] = "--force-dark-mode" + _webengine_flags("--force-dark-mode") else: palette = self.style().standardPalette() plt.style.use("default") - os.environ["QTWEBENGINE_CHROMIUM_FLAGS"] = "" + _webengine_flags(drop=("--force-dark-mode",)) self.setPalette(palette) + self.theme_changed.emit() + + def apply_color_scheme(self, hint) -> None: + """Set Qt color scheme and refresh matplotlib / WebEngine styling.""" + hints = self.styleHints() + hints.blockSignals(True) + try: + hints.setColorScheme(hint) + finally: + hints.blockSignals(False) + self.check_palette(hints.colorScheme()) @property def main_window(self) -> QtWidgets.QMainWindow: diff --git a/activity_browser/ui/dialogs/uncertainty_dialog.py b/activity_browser/ui/dialogs/uncertainty_dialog.py index d0191c794..280a7e527 100644 --- a/activity_browser/ui/dialogs/uncertainty_dialog.py +++ b/activity_browser/ui/dialogs/uncertainty_dialog.py @@ -8,125 +8,30 @@ from qtpy import QtCore, QtGui, QtWidgets import stats_arrays as sa -from activity_browser.bwutils.uncertainty import EMPTY_UNCERTAINTY, standard_uncertainty_fields +from activity_browser.bwutils.uncertainty import ( + DISTRIBUTIONS_WITH_CALCULATED_MEAN, + EMPTY_UNCERTAINTY, + OPTIONAL_UNCERTAINTY_FIELDS, + UNCERTAINTY_VALIDATION_N, + prepare_uncertainty_dict, + uncertainty_dict_is_sampleable, + uncertainty_field_name, + uncertainty_reference_value, + uncertainty_statistics_scalar, + standard_uncertainty_fields, + validate_uncertainty_dict, +) +from stats_arrays import UncertaintyBase from activity_browser.ui.widgets import ABPlot from .uncertainty_pdf_preview import PreviewDensity, preview_density -# ``stats_arrays`` validation uses a short random draw. -_UNCERTAINTY_VALIDATION_N = 24 - - -def _try_build_validate_sample( - dist, - info: dict, - n: int, -) -> Tuple[Optional[np.ndarray], Optional[str]]: - """Build params via ``from_dicts``, run ``stats_arrays`` :meth:`validate`, then a short draw. - - Returns ``(array, None)`` on success, or ``(None, message)`` where *message* is the - exception string from ``stats_arrays`` (e.g. :class:`InvalidParamsError`) or NumPy. - """ - try: - array = dist.from_dicts(info) - dist.validate(array) - dist.random_variables(array, n) - return array, None - except Exception as e: - msg = str(e).strip() or e.__class__.__name__ - return None, msg - - -def _uncertainty_dict_is_sampleable(data: dict) -> bool: - """True if ``stats_arrays`` accepts parameters (``validate``) and a short sample works.""" - try: - uc_type = int(data.get("uncertainty type", 0)) - except Exception: - return False - if uc_type < 0 or uc_type >= len(sa.uncertainty_choices): - return False - dist = sa.uncertainty_choices[uc_type] - if dist.id in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id): - return True - arr, err = _try_build_validate_sample(dist, data, _UNCERTAINTY_VALIDATION_N) - return arr is not None - - _MSG_PREVIEW_INCOMPLETE = ( "Some required parameters are missing or not accepted by the validator. " "Complete them to preview the distribution and enable OK." ) -def _scalar_stats_value(x) -> float: - """Single float from ``statistics()`` / similar (may be scalar or 0-d / 1-element array).""" - if x is None: - return float("nan") - arr = np.asarray(x, dtype=float).ravel() - if arr.size == 0: - return float("nan") - return float(arr[0]) - - -def _discrete_uniform_expected_first_row(array: np.ndarray) -> float: - """Expected value for integers in ``[minimum, maximum)``, per :class:`stats_arrays.DiscreteUniform`.""" - if array is None or len(array) == 0: - return float("nan") - row = array[0] - lo = int(round(float(row["minimum"]))) - hi = int(round(float(row["maximum"]))) - if hi < lo: - lo, hi = hi, lo - if hi <= lo: - return float("nan") - return (float(lo) + float(hi) - 1.0) / 2.0 - - -def _uncertainty_dict_is_sampleable(data: dict) -> bool: - """True if ``stats_arrays`` accepts parameters (``validate``) and a short sample works.""" - try: - uc_type = int(data.get("uncertainty type", 0)) - except Exception: - return False - if uc_type < 0 or uc_type >= len(sa.uncertainty_choices): - return False - dist = sa.uncertainty_choices[uc_type] - if dist.id in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id): - return True - arr, err = _try_build_validate_sample(dist, data, _UNCERTAINTY_VALIDATION_N) - return arr is not None - - -_MSG_PREVIEW_INCOMPLETE = ( - "Some required parameters are missing or not accepted by the validator. " - "Complete them to preview the distribution and enable OK." -) - - -def _scalar_stats_value(x) -> float: - """Single float from ``statistics()`` / similar (may be scalar or 0-d / 1-element array).""" - if x is None: - return float("nan") - arr = np.asarray(x, dtype=float).ravel() - if arr.size == 0: - return float("nan") - return float(arr[0]) - - -def _discrete_uniform_expected_first_row(array: np.ndarray) -> float: - """Expected value for integers in ``[minimum, maximum)``, per :class:`stats_arrays.DiscreteUniform`.""" - if array is None or len(array) == 0: - return float("nan") - row = array[0] - lo = int(round(float(row["minimum"]))) - hi = int(round(float(row["maximum"]))) - if hi < lo: - lo, hi = hi, lo - if hi <= lo: - return float("nan") - return (float(lo) + float(hi) - 1.0) / 2.0 - - class UncertaintyDialog(QtWidgets.QDialog): """Single-step dialog for defining a stats_arrays uncertainty. @@ -150,12 +55,6 @@ def __init__(self, parent=None, initial: Optional[dict] = None, *, read_only: bo self.result_array = None # Filled on accept self.result_dict = None # Filled on accept self.previous_dist_id: Optional[int] = None - self.mean_is_calculated = { - sa.TriangularUncertainty.id, - sa.UniformUncertainty.id, - sa.DiscreteUniform.id, - sa.BetaUncertainty.id, - } # Top: distribution selection box1 = QtWidgets.QGroupBox("Select the uncertainty distribution") @@ -331,7 +230,7 @@ def _apply_initial(self, initial: dict) -> None: data = {k: v for k, v in EMPTY_UNCERTAINTY.items()} data.update(initial or {}) # Do not load numerics that cannot be sampled (e.g. Student's T with df <= 0). - if not _uncertainty_dict_is_sampleable(data): + if not uncertainty_dict_is_sampleable(data): try: uc_type = int(data.get("uncertainty type", 0)) except Exception: @@ -348,7 +247,9 @@ def _apply_initial(self, initial: dict) -> None: self.distribution.setCurrentIndex(uc_type) # Fields (string form for QLineEdit) def to_str(val): - return "nan" if val is None or (isinstance(val, float) and np.isnan(val)) else str(val) + if val is None or (isinstance(val, float) and np.isnan(val)): + return "" + return str(val) self.loc.setText(to_str(data.get("loc", np.nan))) self.scale.setText(to_str(data.get("scale", np.nan))) @@ -369,129 +270,45 @@ def _apply_read_only_mode(self) -> None: cancel_btn.setDefault(True) @property - def _distribution_loc_label(self) -> str: - if self.dist is None: - return "Mean / location:" - if self.dist.id == sa.BernoulliUncertainty.id: - return "Probability (0 ≤ p ≤ 1):" - if self.dist.id == sa.LognormalUncertainty.id: - return "Loc (ln(mean)):" - elif self.dist.id == sa.TriangularUncertainty.id: - return "Mode:" - elif self.dist.id == sa.BetaUncertainty.id: - return "Alpha (α):" - elif self.dist.id in {sa.GammaUncertainty.id, sa.WeibullUncertainty.id}: - return "Loc / offset:" - else: - return "Mean / location:" + def _field_widgets(self) -> dict[str, tuple[QtWidgets.QWidget, QtWidgets.QLineEdit]]: + return { + "loc": (self.loc_label, self.loc), + "scale": (self.scale_label, self.scale), + "shape": (self.shape_label, self.shape), + "minimum": (self.min_label, self.minimum), + "maximum": (self.max_label, self.maximum), + } - def _refresh_axis_labels(self) -> None: - """Scale/shape captions aligned with stats_arrays parameter names.""" - if self.dist is None: - return - d = self.dist.id - if d in ( - sa.NormalUncertainty.id, - sa.LognormalUncertainty.id, - sa.StudentsTUncertainty.id, - sa.GeneralizedExtremeValueUncertainty.id, - ): - self.scale_label.setText("Scale (σ):") - elif d == sa.GammaUncertainty.id: - self.scale_label.setText("Scale (θ):") - elif d == sa.WeibullUncertainty.id: - self.scale_label.setText("Scale (λ):") - else: - self.scale_label.setText("Sigma/scale:") - if d == sa.BetaUncertainty.id: - self.shape_label.setText("Beta (β):") - elif d == sa.StudentsTUncertainty.id: - self.shape_label.setText("Degrees of freedom (ν):") - elif d == sa.GammaUncertainty.id: - self.shape_label.setText("Shape (k):") - elif d == sa.WeibullUncertainty.id: - self.shape_label.setText("Shape (k):") - else: - self.shape_label.setText("Shape:") - - def _hide_params(self, *params, hide: bool = True) -> None: - if "loc" in params: - self.loc_label.setHidden(hide) - self.loc.setHidden(hide) - if "scale" in params: - self.scale_label.setHidden(hide) - self.scale.setHidden(hide) - if "shape" in params: - self.shape_label.setHidden(hide) - self.shape.setHidden(hide) - if "min" in params: - self.min_label.setHidden(hide) - self.minimum.setHidden(hide) - if "max" in params: - self.max_label.setHidden(hide) - self.maximum.setHidden(hide) + def _sync_parameter_fields(self) -> None: + """Show/hide inputs and captions from ``standard_uncertainty_fields``.""" + dist_id = self.dist.id + active = set(standard_uncertainty_fields(dist_id)) + for key, (label, widget) in self._field_widgets.items(): + visible = key in active + label.setVisible(visible) + widget.setVisible(visible) + if visible: + label.setText(f"{uncertainty_field_name(dist_id, key)}:") + if widget.text().strip().lower() == "nan": + widget.clear() + self.neg_samples_cb.setVisible( + dist_id in (sa.GammaUncertainty.id, sa.WeibullUncertainty.id) + ) def _on_distribution_changed(self, index: int) -> None: + self._plot_refresh_timer.stop() self.dist = sa.uncertainty_choices[index] + self._sync_parameter_fields() - # Show/hide fields per stats_arrays parameter usage - if self.dist.id in {0, 1}: # Undefined / NoUncertainty - self._hide_params("loc", "scale", "shape", "min", "max") - self.neg_samples_cb.setHidden(True) - elif self.dist.id in {2, 3}: # Lognormal / Normal - self._hide_params("shape", "min", "max") - self._hide_params("loc", "scale", hide=False) - self.neg_samples_cb.setHidden(True) - elif self.dist.id in {4, 7}: # Uniform / Discrete uniform - self._hide_params("loc", "scale", "shape") - self._hide_params("min", "max", hide=False) - self.neg_samples_cb.setHidden(True) - if self.dist.id == sa.UniformUncertainty.id: - self.min_label.setText("Minimum:") - self.max_label.setText("Maximum:") - else: - # ``stats_arrays.DiscreteUniform``: integers in [minimum, maximum). - self.min_label.setText("Minimum (inclusive):") - self.max_label.setText("Maximum (exclusive):") - elif self.dist.id == sa.TriangularUncertainty.id: # Triangular - self._hide_params("scale", "shape") - self._hide_params("loc", "min", "max", hide=False) - self.neg_samples_cb.setHidden(True) - elif self.dist.id == sa.BernoulliUncertainty.id: # Bernoulli — loc only (probability) - self._hide_params("scale", "shape", "min", "max") - self._hide_params("loc", hide=False) - self.neg_samples_cb.setHidden(True) - elif self.dist.id == sa.BetaUncertainty.id: # Beta — loc (α), shape (β), optional bounds; no ``scale`` - self._hide_params("scale") - self._hide_params("loc", "shape", "min", "max", hide=False) - self.neg_samples_cb.setHidden(True) - elif self.dist.id == sa.GeneralizedExtremeValueUncertainty.id: # GEV — μ, σ only (ξ must be 0) - self._hide_params("shape", "min", "max") - self._hide_params("loc", "scale", hide=False) - self.neg_samples_cb.setHidden(True) - elif self.dist.id in ( - sa.WeibullUncertainty.id, - sa.GammaUncertainty.id, - sa.StudentsTUncertainty.id, - ): - self._hide_params("min", "max") - self._hide_params("loc", "scale", "shape", hide=False) - self.neg_samples_cb.setHidden( - self.dist.id not in (sa.WeibullUncertainty.id, sa.GammaUncertainty.id) - ) - - # Special handling (lognormal and calculated mean label) if self.dist.id == sa.LognormalUncertainty.id: self.mean.setHidden(False) self.mean_label.setHidden(False) - # Convert existing loc to log-space if coming from non-lognormal if self.previous_dist_id is not None and self.previous_dist_id != sa.LognormalUncertainty.id: self._extract_lognormal_loc_from_mean() self._sync_mean_from_loc() else: self.mean.setHidden(True) self.mean_label.setHidden(True) - # If switching away from lognormal, set loc to linear amount if mean present if self.previous_dist_id == sa.LognormalUncertainty.id: try: mean_val = float(self.mean.text()) if self.mean.text() else np.nan @@ -500,20 +317,12 @@ def _on_distribution_changed(self, index: int) -> None: except Exception: pass - # Calculated mean visibility - show_calc = self.dist.id in self.mean_is_calculated + show_calc = self.dist.id in DISTRIBUTIONS_WITH_CALCULATED_MEAN self.calc_mean_label.setHidden(not show_calc) self.calc_mean.setHidden(not show_calc) - - # No parameter inputs for undefined / no uncertainty — hide entire group self.fields_box.setVisible( - self.dist.id - not in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id) + self.dist.id not in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id) ) - - # Update labels - self.loc_label.setText(self._distribution_loc_label) - self._refresh_axis_labels() self.previous_dist_id = self.dist.id self._generate_plot() @@ -567,86 +376,35 @@ def _check_negative(self) -> None: val = float("nan") self.negative.setChecked(bool(not np.isnan(val) and val < 0)) - def _standard_dist_fields(self, dist_id: int) -> list: - return standard_uncertainty_fields(dist_id) - - def _completed_active_fields(self) -> bool: - dist_id = self.dist.id - - def ok_lineedit(le: QtWidgets.QLineEdit) -> bool: - return bool(le.hasAcceptableInput() and le.text()) - - if dist_id in (0, 1): - return True - if dist_id in (sa.LognormalUncertainty.id, sa.NormalUncertainty.id): - return ok_lineedit(self.loc) and ok_lineedit(self.scale) - if dist_id == sa.UniformUncertainty.id: - return ok_lineedit(self.minimum) and ok_lineedit(self.maximum) - if dist_id == sa.DiscreteUniform.id: - if not ok_lineedit(self.maximum): - return False - if not self.minimum.text().strip(): + def _field_has_value(self, text: str) -> bool: + """False for blank fields and the ``nan`` placeholder shown in line edits.""" + t = text.strip().lower() + return bool(t and t != "nan") + + def _any_parameter_text(self) -> bool: + """True if the user has entered any value for the active distribution.""" + optional = OPTIONAL_UNCERTAINTY_FIELDS.get(self.dist.id, frozenset()) + for field in standard_uncertainty_fields(self.dist.id): + if field in optional: + continue + _, widget = self._field_widgets[field] + if widget.isVisible() and self._field_has_value(widget.text()): return True - return ok_lineedit(self.minimum) - if dist_id == sa.TriangularUncertainty.id: - if not ( - ok_lineedit(self.minimum) - and ok_lineedit(self.maximum) - and ok_lineedit(self.loc) - ): - return False - try: - return ( - float(self.minimum.text()) - < float(self.loc.text()) - < float(self.maximum.text()) - ) - except Exception: - return False - if dist_id == sa.BernoulliUncertainty.id: - if not ok_lineedit(self.loc): - return False - try: - p = float(self.loc.text()) - return 0.0 <= p <= 1.0 - except Exception: - return False - if dist_id in ( - sa.WeibullUncertainty.id, - sa.GammaUncertainty.id, - sa.StudentsTUncertainty.id, - ): - return ( - ok_lineedit(self.loc) - and ok_lineedit(self.scale) - and ok_lineedit(self.shape) - ) - if dist_id == sa.BetaUncertainty.id: - if not (ok_lineedit(self.loc) and ok_lineedit(self.shape)): - return False - try: - if float(self.loc.text()) <= 0 or float(self.shape.text()) <= 0: - return False - except Exception: - return False - if not self.minimum.text().strip() and not self.maximum.text().strip(): - return True - if not (ok_lineedit(self.minimum) and ok_lineedit(self.maximum)): - return False - try: - return float(self.minimum.text()) < float(self.maximum.text()) - except Exception: - return False - if dist_id == sa.GeneralizedExtremeValueUncertainty.id: - return ok_lineedit(self.loc) and ok_lineedit(self.scale) return False + def _parse_field(self, text: str) -> float: + if not self._field_has_value(text): + return float("nan") + try: + return float(text) + except (TypeError, ValueError): + return float("nan") + @property def _uncertainty_info(self) -> dict: - data = {k: v for k, v in EMPTY_UNCERTAINTY.items()} - data["uncertainty type"] = self.distribution.currentIndex() if self.dist is None: - return data + return {**EMPTY_UNCERTAINTY} + data = {**EMPTY_UNCERTAINTY, "uncertainty type": self.dist.id} if self.dist.id == sa.LognormalUncertainty.id: data["negative"] = bool(self.negative.isChecked()) elif self.dist.id in (sa.GammaUncertainty.id, sa.WeibullUncertainty.id): @@ -654,59 +412,29 @@ def _uncertainty_info(self) -> dict: else: data["negative"] = False - def as_float(txt: str) -> float: - try: - val = float(txt) - return val - except Exception: - return float("nan") - - for field in self._standard_dist_fields(data["uncertainty type"]): - widget = { - "loc": self.loc, - "scale": self.scale, - "shape": self.shape, - "minimum": self.minimum, - "maximum": self.maximum, - }[field] - data[field] = as_float(widget.text()) - # stats_arrays GEV implementation only supports xi (shape) == 0 - if self.dist.id == sa.GeneralizedExtremeValueUncertainty.id: - data["shape"] = 0.0 - return data + for field in standard_uncertainty_fields(self.dist.id): + _, widget = self._field_widgets[field] + data[field] = self._parse_field(widget.text()) + return prepare_uncertainty_dict(data, self.dist) def _structured_array_if_sampleable(self) -> Tuple[Optional[np.ndarray], Optional[str]]: - """Build params, ``stats_arrays`` validation, and a short random draw.""" - return _try_build_validate_sample(self.dist, self._uncertainty_info, _UNCERTAINTY_VALIDATION_N) + return validate_uncertainty_dict(self._uncertainty_info, self.dist, UNCERTAINTY_VALIDATION_N) def _ok_enabled(self) -> bool: if self.dist is None: return False if self.dist.id in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id): return True - if not self._completed_active_fields(): - return False - arr, _ = self._structured_array_if_sampleable() - return arr is not None + array, _ = self._structured_array_if_sampleable() + return array is not None def _update_ok_state(self, structured: Optional[np.ndarray] = None) -> None: if self._read_only: return ok_btn = self.buttons.button(QtWidgets.QDialogButtonBox.Ok) - if self.dist is None: - ok_btn.setEnabled(False) - return - if self.dist.id in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id): - ok_btn.setEnabled(True) - return - if not self._completed_active_fields(): - ok_btn.setEnabled(False) - return - if structured is not None: - ok_btn.setEnabled(True) - return - arr, _ = self._structured_array_if_sampleable() - ok_btn.setEnabled(arr is not None) + ok_btn.setEnabled( + structured is not None if structured is not None else self._ok_enabled() + ) def _schedule_plot_refresh(self) -> None: """Refresh OK state immediately; defer heavy matplotlib work to avoid UI stalls.""" @@ -754,54 +482,24 @@ def _generate_plot(self) -> None: self._update_ok_state() return - if not self._completed_active_fields(): - self._hide_plot_preview(_MSG_PREVIEW_INCOMPLETE) - self._update_ok_state() - return - array, sample_err = self._structured_array_if_sampleable() if array is None: - self._hide_plot_preview( - sample_err or "Invalid parameters for this distribution." + msg = ( + _MSG_PREVIEW_INCOMPLETE + if not self._any_parameter_text() + else (sample_err or "Invalid parameters for this distribution.") ) + self._hide_plot_preview(msg) self._update_ok_state() return - if self.dist.id in self.mean_is_calculated: - try: - if self.dist.id == sa.DiscreteUniform.id: - calc = _discrete_uniform_expected_first_row(array) - self.calc_mean.setText( - str(float(calc)) if np.isfinite(calc) else "nan" - ) - else: - try: - calc = self.dist.statistics(array).get("mean") - except TypeError: - array = self.dist.fix_nan_minimum(array) - calc = (array["maximum"] + array["minimum"]) / 2 - calc = calc.mean() if isinstance(calc, np.ndarray) else calc - self.calc_mean.setText(str(float(calc))) - except Exception: - self.calc_mean.setText("nan") + if self.dist.id in DISTRIBUTIONS_WITH_CALCULATED_MEAN: + calc = uncertainty_statistics_scalar(self.dist, array, "mean") + self.calc_mean.setText(str(calc) if np.isfinite(calc) else "nan") - try: - if self.dist.id == sa.LognormalUncertainty.id: - ref_lin = _scalar_stats_value( - self.dist.statistics(array).get("median") - ) - elif self.dist.id == sa.DiscreteUniform.id: - ref_lin = _discrete_uniform_expected_first_row(array) - else: - ref_lin = _scalar_stats_value( - self.dist.statistics(array).get("mean") - ) - except Exception as e: - logger.debug( - "Uncertainty preview skipped (invalid statistics): {}", - e, - ) - self._hide_plot_preview(str(e).strip() or e.__class__.__name__) + ref_lin = uncertainty_reference_value(self.dist, array) + if not np.isfinite(ref_lin): + self._hide_plot_preview("Could not compute distribution statistics for preview.") self._update_ok_state() return @@ -816,7 +514,7 @@ def _generate_plot(self) -> None: try: self._plot_message.hide() self._plot_message.clear() - self.plot.plot_analytical(curve, ref_lin) + self.plot.plot_analytical(curve, ref_lin, title=self.dist.description) except Exception as e: logger.warning("Uncertainty preview could not be drawn: {}", e) self._hide_plot_preview(str(e).strip() or e.__class__.__name__) @@ -833,7 +531,7 @@ def _on_accept(self) -> None: return try: self.result_dict = self._uncertainty_info - self.result_array = self.dist.from_dicts(self._uncertainty_info) + self.result_array = UncertaintyBase.from_dicts(self._uncertainty_info) except Exception as e: QtWidgets.QMessageBox.warning( self, @@ -862,16 +560,16 @@ def __init__(self, parent=None): m = lay.contentsMargins() lay.setContentsMargins(m.left(), m.top(), m.right(), 6) - def plot_analytical(self, curve: PreviewDensity, vline_x: float) -> None: + def plot_analytical( + self, curve: PreviewDensity, vline_x: float, *, title: str = "" + ) -> None: """Plot ``stats_arrays`` / SciPy PDF or PMF (no random sampling).""" self.setMinimumHeight(348) self.setMaximumHeight(16777215) - exp = QtWidgets.QSizePolicy.Policy.Expanding - self.setSizePolicy(exp, exp) - self.canvas.setSizePolicy(exp, exp) + self.setVisible(True) self.reset_plot() - # Match figure size to widget before plotting so axes bbox is non-zero. - self.sync_figure_to_widget() + if title: + self.ax.set_title(title, fontsize=10, pad=6) if curve.kind == "bar": self.ax.set_ylabel("PMF", labelpad=6) @@ -897,7 +595,7 @@ def plot_analytical(self, curve: PreviewDensity, vline_x: float) -> None: ) self.ax.set_xlabel(curve.xlabel, labelpad=7) - self._set_plot_chrome_white() + self._sync_plot_to_theme() try: if np.isfinite(vline_x): self.ax.axvline(vline_x, label=curve.vline_legend, c="r", ymax=0.98) @@ -912,6 +610,17 @@ def plot_analytical(self, curve: PreviewDensity, vline_x: float) -> None: self.ax.tick_params(axis="both", which="major", labelsize=9) self.ax.tick_params(axis="x", which="major", pad=4) + # Layout after Qt assigns the restored height (first show is often still 0×0). + QtCore.QTimer.singleShot(0, self._fit_preview_to_widget) + + def _fit_preview_to_widget(self, _attempt: int = 0) -> None: + self.sync_figure_to_widget() + if not self._canvas_has_size(): + if _attempt < 8: + QtCore.QTimer.singleShot(0, lambda: self._fit_preview_to_widget(_attempt + 1)) + else: + self.canvas.draw_idle() + return try: self.figure.tight_layout(pad=0.28, h_pad=0.55, w_pad=0.35) except Exception: @@ -919,7 +628,6 @@ def plot_analytical(self, curve: PreviewDensity, vline_x: float) -> None: else: p = self.figure.subplotpars h_in = float(self.figure.get_figheight()) - # Slightly more bottom margin when the figure is very short (labels + ticks). bottom_floor = max(0.14, min(0.24, 1.05 / max(h_in, 0.25))) self.figure.subplots_adjust( left=max(p.left, 0.07), @@ -927,10 +635,7 @@ def plot_analytical(self, curve: PreviewDensity, vline_x: float) -> None: top=min(p.top, 0.97), bottom=max(p.bottom, bottom_floor), ) - - self.sync_figure_to_widget() self.canvas.draw_idle() - self.setVisible(True) __all__ = ["UncertaintyDialog"] diff --git a/activity_browser/ui/dialogs/uncertainty_pdf_preview.py b/activity_browser/ui/dialogs/uncertainty_pdf_preview.py index b92228c1f..bf2503bfe 100644 --- a/activity_browser/ui/dialogs/uncertainty_pdf_preview.py +++ b/activity_browser/ui/dialogs/uncertainty_pdf_preview.py @@ -1,5 +1,5 @@ # -*- coding: utf-8 -*- -"""Analytical PDF/PMF curves for uncertainty dialog preview (no Monte Carlo histogram).""" +"""PDF/PMF preview: ``stats_arrays`` :meth:`pdf` first, SciPy only where missing.""" from __future__ import annotations from typing import NamedTuple, Optional @@ -8,56 +8,19 @@ import scipy.stats as st import stats_arrays as sa +from activity_browser.bwutils.uncertainty import as_scalar -class PreviewDensity(NamedTuple): - """Curve data for :class:`SimpleDistributionPlot`.""" +class PreviewDensity(NamedTuple): x: np.ndarray y: np.ndarray - # "line" = continuous density; "bar" = PMF at discrete support kind: str = "line" xlabel: str = "Value" vline_legend: str = "Mean / amount" -def _row(arr: np.ndarray) -> np.void: - return arr[0] - - -def _f(row: np.void, key: str, default: float = np.nan) -> float: - v = row[key] - if isinstance(v, np.ndarray): - v = v.flat[0] - try: - return float(v) - except (TypeError, ValueError): - return float("nan") - - -def _bool(row: np.void, key: str) -> bool: - return bool(row[key]) - - -def _try_stats_arrays_pdf(dist, arr: np.ndarray, xs: Optional[np.ndarray]) -> Optional[tuple[np.ndarray, np.ndarray]]: - try: - x, y = dist.pdf(arr, xs) - except (NotImplementedError, TypeError, ValueError): - return None - x = np.asarray(x, dtype=float).ravel() - y = np.asarray(y, dtype=float).ravel() - if x.size == 0 or y.size == 0 or x.size != y.size: - return None - m = np.nanmax(y) - if not np.isfinite(m) or m <= 1e-30: - return None - ok = np.isfinite(x) & np.isfinite(y) - if not ok.any(): - return None - return x[ok], y[ok] - - -def _linspace_ppf(rv: st.rv_continuous, n: int, lo: float = 0.001, hi: float = 0.999) -> np.ndarray: - a, b = rv.ppf(lo), rv.ppf(hi) +def _linspace_ppf(rv: st.rv_continuous, n: int) -> np.ndarray: + a, b = rv.ppf(0.001), rv.ppf(0.999) if not np.isfinite(a): a = rv.ppf(0.02) if not np.isfinite(b): @@ -67,148 +30,107 @@ def _linspace_ppf(rv: st.rv_continuous, n: int, lo: float = 0.001, hi: float = 0 return np.linspace(a, b, n) -def _lognormal_positive_curve(row: np.void, n: int) -> PreviewDensity: - """SciPy log-normal PDF (stats_arrays ``pdf`` can return zeros for some ``loc`` values).""" - mu = _f(row, "loc") - sig = _f(row, "scale") - rv = st.lognorm(s=sig, scale=np.exp(mu)) - xp = _linspace_ppf(rv, n) - yp = rv.pdf(xp) - return PreviewDensity(x=xp, y=yp, kind="line", xlabel="Value", vline_legend="Median") - - -def _lognormal_negative_curve(row: np.void, n: int) -> PreviewDensity: - mu = _f(row, "loc") - sig = _f(row, "scale") - rv = st.lognorm(s=sig, scale=np.exp(mu)) - xp = _linspace_ppf(rv, max(n // 2, 50)) - yp = rv.pdf(xp) - xn = -xp[::-1] - yn = yp[::-1] - x = np.concatenate([xn[:-1], xp]) - y = np.concatenate([yn[:-1], yp]) - return PreviewDensity(x=x, y=y, kind="line", xlabel="Value", vline_legend="Median") - - -def _gamma_weibull_scipy_pdf(dist_id: int, row: np.void, n: int) -> PreviewDensity: - loc = _f(row, "loc") - if not np.isfinite(loc): - loc = 0.0 - scale = _f(row, "scale") - shape = _f(row, "shape") - neg = _bool(row, "negative") - if dist_id == sa.GammaUncertainty.id: - rv = st.gamma(a=shape, loc=loc, scale=scale) - else: - rv = st.weibull_min(c=shape, loc=loc, scale=scale) - xp = _linspace_ppf(rv, max(n // 2, 50)) - yp = rv.pdf(xp) - if neg: - xn = -xp[::-1] - yn = yp[::-1] - x = np.concatenate([xn[:-1], xp]) - y = np.concatenate([yn[:-1], yp]) - else: - x, y = xp, yp - vleg = "Mean / amount" - return PreviewDensity(x=x, y=y, kind="line", xlabel="Value", vline_legend=vleg) - - -def _students_t_curve(row: np.void, n: int) -> PreviewDensity: - df = _f(row, "shape") - loc = _f(row, "loc") - if not np.isfinite(loc): - loc = 0.0 - sc = _f(row, "scale") - if not np.isfinite(sc) or sc <= 0: - sc = 1.0 - rv = st.t(df=df, loc=loc, scale=sc) - xs = _linspace_ppf(rv, n) - ys = rv.pdf(xs) - return PreviewDensity(x=xs, y=ys, kind="line", xlabel="Value", vline_legend="Mean / amount") - - -def _gumbel_curve(row: np.void, n: int) -> PreviewDensity: - # stats_arrays GEV only supports xi=0 and draws from gumbel(loc, scale) - loc = _f(row, "loc") - sc = _f(row, "scale") - rv = st.gumbel_r(loc=loc, scale=sc) - xs = _linspace_ppf(rv, n) - ys = rv.pdf(xs) - return PreviewDensity(x=xs, y=ys, kind="line", xlabel="Value", vline_legend="Mean / amount") - - -def _bernoulli_pmf(row: np.void) -> PreviewDensity: - p = _f(row, "loc") - p = min(max(p, 0.0), 1.0) - x = np.array([0.0, 1.0], dtype=float) - y = np.array([1.0 - p, p], dtype=float) - return PreviewDensity(x=x, y=y, kind="bar", xlabel="Value", vline_legend="Probability") - - -def _discrete_uniform_pmf(row: np.void) -> PreviewDensity: - """PMF on integers ``minimum, minimum+1, …, maximum - 1`` (``maximum`` is exclusive). - - Matches :class:`stats_arrays.DiscreteUniform` / ``numpy.random.randint(low, high)``. - """ - lo = int(round(_f(row, "minimum"))) - hi = int(round(_f(row, "maximum"))) - if hi < lo: - lo, hi = hi, lo - x = np.arange(lo, hi, dtype=float) - if x.size == 0: +def _mirror_negative(xp: np.ndarray, yp: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + xn, yn = -xp[::-1], yp[::-1] + return np.concatenate([xn[:-1], xp]), np.concatenate([yn[:-1], yp]) + + +def _scipy_preview(dist_id: int, row: np.void, n: int) -> Optional[PreviewDensity]: + """SciPy curves for distributions without a working ``stats_arrays`` ``pdf()``.""" + if dist_id == sa.LognormalUncertainty.id: + mu, sig = as_scalar(row["loc"]), as_scalar(row["scale"]) + neg = bool(row["negative"]) + rv = st.lognorm(s=sig, scale=np.exp(mu)) + xp = _linspace_ppf(rv, max(n // 2, 50) if neg else n) + yp = rv.pdf(xp) + if neg: + xp, yp = _mirror_negative(xp, yp) + return PreviewDensity(x=xp, y=yp, kind="line", vline_legend="Median") + + if dist_id in (sa.GammaUncertainty.id, sa.WeibullUncertainty.id): + loc = as_scalar(row["loc"]) + if not np.isfinite(loc): + loc = 0.0 + scale, shape = as_scalar(row["scale"]), as_scalar(row["shape"]) + if dist_id == sa.GammaUncertainty.id: + rv = st.gamma(a=shape, loc=loc, scale=scale) + else: + rv = st.weibull_min(c=shape, loc=loc, scale=scale) + xp = _linspace_ppf(rv, max(n // 2, 50)) + yp = rv.pdf(xp) + if bool(row["negative"]): + xp, yp = _mirror_negative(xp, yp) + return PreviewDensity(x=xp, y=yp, kind="line") + + if dist_id == sa.StudentsTUncertainty.id: + loc = as_scalar(row["loc"]) + if not np.isfinite(loc): + loc = 0.0 + sc = as_scalar(row["scale"]) + if not np.isfinite(sc) or sc <= 0: + sc = 1.0 + rv = st.t(df=as_scalar(row["shape"]), loc=loc, scale=sc) + xs = _linspace_ppf(rv, n) + return PreviewDensity(x=xs, y=rv.pdf(xs), kind="line") + + if dist_id == sa.GeneralizedExtremeValueUncertainty.id: + rv = st.gumbel_r(loc=as_scalar(row["loc"]), scale=as_scalar(row["scale"])) + xs = _linspace_ppf(rv, n) + return PreviewDensity(x=xs, y=rv.pdf(xs), kind="line") + + if dist_id == sa.BernoulliUncertainty.id: + p = min(max(as_scalar(row["loc"]), 0.0), 1.0) + return PreviewDensity( + x=np.array([0.0, 1.0]), + y=np.array([1.0 - p, p]), + kind="bar", + vline_legend="Probability", + ) + + if dist_id == sa.DiscreteUniform.id: + lo = int(round(as_scalar(row["minimum"]))) + hi = int(round(as_scalar(row["maximum"]))) + if hi < lo: + lo, hi = hi, lo + x = np.arange(lo, hi, dtype=float) + if x.size == 0: + return PreviewDensity( + x=np.array([float(lo)]), + y=np.array([0.0]), + kind="bar", + vline_legend="Expected value", + ) return PreviewDensity( - x=np.array([float(lo)]), - y=np.array([0.0]), + x=x, + y=np.full(x.shape, 1.0 / x.size), kind="bar", - xlabel="Value", vline_legend="Expected value", ) - p = 1.0 / float(x.size) - y = np.full_like(x, p) - return PreviewDensity(x=x, y=y, kind="bar", xlabel="Value", vline_legend="Expected value") + return None def preview_density(dist, structured_array: np.ndarray, n_points: int = 400) -> Optional[PreviewDensity]: - """Return density/PMF for the first row of *structured_array*, or None if not drawable.""" - if structured_array is None or len(structured_array) == 0: + """Density or PMF for *dist* (authoritative type) and the first parameter row.""" + if dist is None or structured_array is None or len(structured_array) == 0: return None - row = _row(structured_array) - dist_id = int(row["uncertainty_type"]) - + dist_id = dist.id if dist_id in (sa.UndefinedUncertainty.id, sa.NoUncertainty.id): return None - # Lognormal + negative: stats_arrays pdf() is not usable; mirror positive lognormal. - if dist_id == sa.LognormalUncertainty.id and _bool(row, "negative"): - return _lognormal_negative_curve(row, n_points) + row = structured_array[0] - # Discrete laws: ``stats_arrays`` :meth:`pdf` may return a sampled continuous curve; - # always show the exact PMF as bars. - if dist_id == sa.BernoulliUncertainty.id: - return _bernoulli_pmf(row) - if dist_id == sa.DiscreteUniform.id: - return _discrete_uniform_pmf(row) - - res = _try_stats_arrays_pdf(dist, structured_array, None) - if dist_id == sa.LognormalUncertainty.id and not _bool(row, "negative"): - if res is None or float(np.nanmax(res[1])) <= 1e-30: - return _lognormal_positive_curve(row, n_points) - if res is not None: - x, y = res - vleg = "Median" if dist_id == sa.LognormalUncertainty.id else "Mean / amount" - return PreviewDensity(x=x, y=y, kind="line", xlabel="Value", vline_legend=vleg) - - if dist_id == sa.WeibullUncertainty.id: - return _gamma_weibull_scipy_pdf(dist_id, row, n_points) - if dist_id == sa.GammaUncertainty.id: - return _gamma_weibull_scipy_pdf(dist_id, row, n_points) - if dist_id == sa.StudentsTUncertainty.id: - return _students_t_curve(row, n_points) - if dist_id == sa.GeneralizedExtremeValueUncertainty.id: - return _gumbel_curve(row, n_points) + try: + x, y = dist.pdf(structured_array, None) + x = np.asarray(x, dtype=float).ravel() + y = np.asarray(y, dtype=float).ravel() + if x.size and x.size == y.size and np.isfinite(y).any() and np.nanmax(y) > 1e-30: + ok = np.isfinite(x) & np.isfinite(y) + vleg = "Median" if dist_id == sa.LognormalUncertainty.id else "Mean / amount" + return PreviewDensity(x=x[ok], y=y[ok], kind="line", vline_legend=vleg) + except (NotImplementedError, TypeError, ValueError): + pass - return None + return _scipy_preview(dist_id, row, n_points) __all__ = ["PreviewDensity", "preview_density"] diff --git a/activity_browser/ui/widgets/plot.py b/activity_browser/ui/widgets/plot.py index b7b6eac84..b400cbde3 100644 --- a/activity_browser/ui/widgets/plot.py +++ b/activity_browser/ui/widgets/plot.py @@ -30,9 +30,22 @@ from activity_browser.bwutils.contribution_labels import is_rest_row _GOLDEN_RATIO = 0.618033988749895 -_BASE_SERIES_PALETTE: tuple[tuple[float, float, float, float], ...] = tuple( - plt.get_cmap("tab10").colors -) + tuple(plt.get_cmap("tab20").colors) +DEFAULT_PLOT_PALETTE = "tab20" +PLOT_PALETTES: tuple[str, ...] = ( # qualitative colormaps from https://matplotlib.org/stable/gallery/color/colormap_reference.html + "Pastel1", + "Pastel2", + "Paired", + "Accent", + "okabe_ito", + "Dark2", + "Set1", + "Set2", + "Set3", + "tab10", + "tab20", + "tab20b", + "tab20c", +) _GSA_TYPE_ORDER = ( "technosphere", "biosphere", @@ -55,6 +68,23 @@ def lca_results_tab_from_widget(widget) -> object | None: return None +def _categorical_palette() -> tuple[tuple[float, float, float, float], ...]: + from activity_browser import app + + name = app.settings["appearance"].get("plot_palette", DEFAULT_PLOT_PALETTE) + if name not in PLOT_PALETTES: + name = DEFAULT_PLOT_PALETTE + for candidate in (name, DEFAULT_PLOT_PALETTE): + try: + return tuple( + (float(r), float(g), float(b), 1.0) + for r, g, b in plt.get_cmap(candidate).colors + ) + except (ValueError, TypeError, AttributeError): + continue + return (0.5, 0.5, 0.5, 1.0), + + class ABFigureCanvas(FigureCanvasQTAgg): """Canvas with zero width hint so scroll areas do not widen the main window.""" @@ -79,16 +109,16 @@ class ABPlot(QtWidgets.QWidget): HORIZONTAL_AXIS_LABEL_WRAP_LENGTH = 40 REST_BAR_COLOR = (0.8, 0.8, 0.8, 1.0) - BAR_EDGE_COLOR = "white" BAR_EDGE_WIDTH = 0.3 # --- series color palette ------------------------------------------------- @classmethod def series_color(cls, index: int) -> tuple[float, float, float, float]: - """RGBA for categorical series ``index`` (30 curated colors, then golden-ratio hues).""" - if index < len(_BASE_SERIES_PALETTE): - return _BASE_SERIES_PALETTE[index] + """RGBA for categorical series ``index`` (palette colors, then golden-ratio hues).""" + palette = _categorical_palette() + if index < len(palette): + return palette[index] hue = (index * _GOLDEN_RATIO) % 1.0 rgb = mcolors.hsv_to_rgb((hue, 0.72, 0.88)) return (float(rgb[0]), float(rgb[1]), float(rgb[2]), 1.0) @@ -150,12 +180,13 @@ def legend_column_width_ratio(self) -> float: @staticmethod def set_signed_value_grid(ax, *, horizontal: bool) -> None: """Zero reference line and dashed grid on the value axis.""" + zero_color = plt.rcParams["axes.edgecolor"] ax.set_axisbelow(True) if horizontal: - ax.axvline(0, color="black", linewidth=0.8, zorder=1) + ax.axvline(0, color=zero_color, linewidth=0.8, zorder=1) ax.grid(axis="x", linestyle="dashed", color="grey", alpha=0.7) else: - ax.axhline(0, color="black", linewidth=0.8, zorder=1) + ax.axhline(0, color=zero_color, linewidth=0.8, zorder=1) ax.grid(axis="y", linestyle="dashed", color="grey", alpha=0.7) def plot_bar_strip( @@ -172,7 +203,7 @@ def plot_bar_strip( kw = dict( label=label, color=[color] * len(heights), - edgecolor=self.BAR_EDGE_COLOR, + edgecolor=plt.rcParams["axes.facecolor"], linewidth=self.BAR_EDGE_WIDTH, ) size = thickness * 0.92 @@ -234,7 +265,10 @@ def __init__(self, parent=None): self._axis_contexts: list[dict] = [] self._current_bar_ctx: dict | None = None - self._set_plot_chrome_white() + self._sync_plot_to_theme() + from activity_browser import app + + app.application.theme_changed.connect(self._on_theme_changed) layout = QtWidgets.QVBoxLayout() layout.setContentsMargins(0, 0, 0, 0) @@ -710,13 +744,80 @@ def clear_hover_tooltip(self) -> None: def plot(self, *args, **kwargs): raise NotImplementedError - def _set_plot_chrome_white(self) -> None: - self.figure.patch.set_facecolor("white") - if self.ax is not None: - self.ax.set_facecolor("white") - bg = "background-color: white;" - self.canvas.setStyleSheet(bg) - self.setStyleSheet(bg) + @staticmethod + def _color_line_artists(artists, color) -> None: + if artists is None: + return + if not isinstance(artists, (list, tuple)): + artists = (artists,) + for artist in artists: + if artist is None: + continue + if hasattr(artist, "set_color"): + artist.set_color(color) + elif hasattr(artist, "set_edgecolor"): + artist.set_edgecolor(color) + + def _sync_axes_to_theme(self, ax) -> None: + from matplotlib.container import ErrorbarContainer + + rc = plt.rcParams + edge, tick = rc["axes.edgecolor"], rc["xtick.color"] + ax.set_facecolor(rc["axes.facecolor"]) + for spine in ax.spines.values(): + spine.set_color(edge) + ax.tick_params(axis="both", colors=tick, labelcolor=tick, which="both") + if ax.xaxis.label: + ax.xaxis.label.set_color(rc["axes.labelcolor"]) + if ax.yaxis.label: + ax.yaxis.label.set_color(rc["axes.labelcolor"]) + if ax.title: + ax.title.set_color(rc["text.color"]) + for line in ax.lines: + x, y = np.asarray(line.get_xdata()), np.asarray(line.get_ydata()) + if x.size >= 2 and y.size >= 2 and ( + (np.allclose(y, 0) and not np.allclose(x, 0)) + or (np.allclose(x, 0) and not np.allclose(y, 0)) + ): + line.set_color(edge) + elif ( + line.get_marker() not in (None, "none", "None") + and str(line.get_linestyle()).lower() == "none" + ): + line.set_markerfacecolor(edge) + line.set_markeredgecolor(edge) + for patch in ax.patches: + patch.set_edgecolor(rc["axes.facecolor"]) + for axis in (ax.xaxis, ax.yaxis): + for gridline in axis.get_gridlines(): + gridline.set_color(rc.get("grid.color", "0.5")) + gridline.set_alpha(rc.get("grid.alpha", 0.7)) + for container in ax.containers: + if isinstance(container, ErrorbarContainer): + self._color_line_artists(container.lines[1], edge) + self._color_line_artists(container.lines[2], edge) + + def _sync_plot_to_theme(self) -> None: + rc = plt.rcParams + self.figure.patch.set_facecolor(rc["figure.facecolor"]) + axes = list(self.figure.axes) + if self.ax is not None and self.ax not in axes: + axes.append(self.ax) + for ax in axes: + self._sync_axes_to_theme(ax) + for legend in self.figure.legends: + for text in legend.get_texts(): + text.set_color(rc["text.color"]) + qapp = QtWidgets.QApplication.instance() + if qapp is not None: + bg = f"background-color: {qapp.palette().color(QtGui.QPalette.ColorRole.Window).name()};" + self.canvas.setStyleSheet(bg) + self.setStyleSheet(bg) + + def _on_theme_changed(self) -> None: + self._sync_plot_to_theme() + if self.figure.axes: + self.canvas.draw_idle() def reset_plot(self) -> None: self.clear_hover_tooltip() @@ -724,7 +825,7 @@ def reset_plot(self) -> None: self._current_bar_ctx = None self.figure.clf() self.ax = self.figure.add_subplot(111) - self._set_plot_chrome_white() + self._sync_plot_to_theme() def add_legend(self, *args, ax=None, **kwargs): kwargs.setdefault("ncol", 1) @@ -747,7 +848,7 @@ def finish_plot( ) -> None: """Apply fonts, draw, and wire hover tooltips. Call once per :meth:`plot`.""" self.apply_standard_fonts() - self._set_plot_chrome_white() + self._sync_plot_to_theme() self.sync_figure_to_widget() if self._canvas_has_size(): self.canvas.draw() diff --git a/tests/actions/test_database_actions.py b/tests/actions/test_database_actions.py index ddcf7694c..d28a11e96 100644 --- a/tests/actions/test_database_actions.py +++ b/tests/actions/test_database_actions.py @@ -113,43 +113,6 @@ def test_fix_broken_groups_removes_orphan_activity_groups(basic_database): assert not {g.name for g in Group.select()} & activity_groups -def test_database_duplicate(monkeypatch, qtbot, basic_database): - from activity_browser.app.actions.database.database_duplicate import NewDatabaseDialog, DuplicateDatabaseDialog - from activity_browser.bwutils.commontasks import count_database_records - - dup_db = "db_that_is_duplicated" - source_count = count_database_records(basic_database.name) - - monkeypatch.setattr( - NewDatabaseDialog, - "get_new_database_data", - staticmethod(lambda *args, **kwargs: (dup_db, "functional_sqlite", True)), - ) - - assert dup_db not in bd.databases - - app.actions.DatabaseDuplicate.run(basic_database.name) - - dialog = app.main_window.findChild(DuplicateDatabaseDialog) - with qtbot.waitSignal(dialog.dup_thread.finished, timeout=60 * 1000): - pass - - assert basic_database.name in bd.databases - assert dup_db in bd.databases - assert count_database_records(dup_db) == source_count - - loader = app.metadata.loader - for _ in range(200): - if len(app.metadata.get_database_metadata(dup_db, ["name"])) == source_count: - break - if loader.secondary_status != "done": - qtbot.wait(50) - continue - qtbot.wait(50) - - assert len(app.metadata.get_database_metadata(dup_db, ["name"])) == source_count - - def test_database_export_excel(monkeypatch, qtbot, basic_database, tmp_path): """Test exporting a database to Excel format.""" from activity_browser.app.actions.database.database_export_excel import ExportExcelSetup diff --git a/tests/conftest.py b/tests/conftest.py index 440080fe4..da4cd089e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,10 +1,10 @@ -from copy import deepcopy from importlib import reload from loguru import logger import pandas as pd import pytest import os +import time import bw2data as bd from PySide6 import QtCore, QtWidgets @@ -21,40 +21,56 @@ os.environ["AB_SKIP_SETTINGS_ON_STARTUP"] = "1" os.environ["AB_NO_SEARCHER"] = "1" +_MAIN_WINDOW_READY = False -def _destroy_main_window(qtbot): - """Close dynamic tabs/panes and reset the main-window singleton between tests.""" - from activity_browser import app - from activity_browser.ui import core - qapp = QtWidgets.QApplication.instance() - mw = getattr(app, "main_window", None) +def _wait_for_loader(qapp, loader, timeout: float = 30.0) -> None: + """Poll metadata loader with short sleeps instead of blocking 1s per iteration.""" + deadline = time.monotonic() + timeout + while loader.secondary_status != "done" and time.monotonic() < deadline: + qapp.processEvents(QtCore.QEventLoop.ProcessEventsFlag.AllEvents) + time.sleep(0.05) + if loader.secondary_status != "done": + raise TimeoutError("Metadata loader did not finish in time.") + + +def _ensure_main_window() -> None: + """Create the main window once per process; tests only reset lightweight state.""" + global _MAIN_WINDOW_READY + from activity_browser import app + from activity_browser.bwutils.metadata import metadata - if mw is not None and core.qt_is_valid(mw): - central = mw.centralWidget() - if central is not None and core.qt_is_valid(central): - for index in range(central.count() - 1, -1, -1): - widget = central.widget(index) - central.removeTab(index) - if widget is not None and core.qt_is_valid(widget): - widget.deleteLater() + if _MAIN_WINDOW_READY: + return - for pane in list(mw.panes()): - if core.qt_is_valid(pane): - pane.hide() - pane.deleteLater() + reload(metadata) + reload(app.main) + reload(app) + _MAIN_WINDOW_READY = True - mw.close() - mw.deleteLater() - if qapp is not None: - for _ in range(3): - qapp.processEvents(QtCore.QEventLoop.ProcessEventsFlag.AllEvents) - qtbot.wait(100) +def _reset_main_window(qtbot) -> None: + """Close extra tabs opened during a test; keep the main window alive.""" + from activity_browser import app + from activity_browser.ui import core - from activity_browser.app.main import MainWindow + qapp = QtWidgets.QApplication.instance() + mw = getattr(app, "main_window", None) + if mw is None or not core.qt_is_valid(mw): + return + + central = mw.centralWidget() + if central is not None and core.qt_is_valid(central): + while central.count() > 1: + index = central.count() - 1 + widget = central.widget(index) + central.removeTab(index) + if widget is not None and core.qt_is_valid(widget): + widget.deleteLater() - MainWindow._instance = None + if qapp is not None: + qapp.processEvents(QtCore.QEventLoop.ProcessEventsFlag.AllEvents) + qtbot.wait(10) @pytest.fixture @@ -64,7 +80,6 @@ def no_exception_dialogs(monkeypatch): monkeypatch.setattr(QtWidgets.QMessageBox, "critical", lambda *args, **kwargs: None) yield - # No need to undo the monkeypatch, pytest does it automatically @pytest.fixture @@ -73,52 +88,29 @@ def main_window(qtbot, monkeypatch, no_exception_dialogs): from activity_browser import app from activity_browser.bwutils.metadata import metadata - _destroy_main_window(qtbot) - - # Reload modules to ensure a clean state for each test - reload(metadata) - reload(app.main) - reload(app) + _ensure_main_window() metadata.dataframe = pd.DataFrame() - app.main_window.show() yield app.main_window - _destroy_main_window(qtbot) - + _reset_main_window(qtbot) @pytest.fixture @bw2test def basic_database(qapp, main_window): - import time from activity_browser.app import metadata from fixtures.basic import CALCULATION_SETUP, DATABASE, METHOD qapp.processEvents(QtCore.QEventLoop.ProcessEventsFlag.AllEvents) - - i = 0 - while metadata.loader.secondary_status != "done" and i < 60: - logger.warning("Waiting for project load to finish") - time.sleep(1) - qapp.processEvents(QtCore.QEventLoop.ProcessEventsFlag.AllEvents) - i += 1 + _wait_for_loader(qapp, metadata.loader) db = write_functional_database("basic", DATABASE, process=False, mark_dirty=True) write_method("basic_method", METHOD, process=False) write_calculation_setup("basic_calculation_setup", CALCULATION_SETUP) - i = 0 - while metadata.loader.secondary_status != "done" and i < 60: - logger.warning("Waiting for database load to finish...") - time.sleep(1) - qapp.processEvents(QtCore.QEventLoop.ProcessEventsFlag.AllEvents) - i += 1 - - if i >= 60: - raise TimeoutError("Metadata loader did not finish in time.") - + _wait_for_loader(qapp, metadata.loader) yield db diff --git a/tests/test_activity_new_elementary_flow.py b/tests/test_activity_new_elementary_flow.py index c1559956b..ec2f2ebc7 100644 --- a/tests/test_activity_new_elementary_flow.py +++ b/tests/test_activity_new_elementary_flow.py @@ -4,7 +4,7 @@ from activity_browser import app from activity_browser.app.actions.activity import new_elementary_flow as mod -from activity_browser.bwutils.commontasks import get_writable_databases, is_node_biosphere +from activity_browser.bwutils.commontasks import is_node_biosphere class _AcceptedDialog: @@ -18,11 +18,6 @@ def get_data(self): return ("custom emission", "kg", "emission", ("air", "custom")) -def test_parse_categories(): - assert mod._parse_categories("") == () - assert mod._parse_categories("air, non-urban") == ("air", "non-urban") - - def _make_database_writable(db_name: str) -> None: import bw2data as bd @@ -30,11 +25,6 @@ def _make_database_writable(db_name: str) -> None: bd.databases.flush() -def test_writable_databases_includes_basic(basic_database): - _make_database_writable(basic_database.name) - assert "basic" in get_writable_databases() - - def test_activity_new_elementary_flow(basic_database, monkeypatch): _make_database_writable(basic_database.name) monkeypatch.setattr(mod, "ElementaryFlowDialog", _AcceptedDialog) diff --git a/tests/test_database_roundtrip.py b/tests/test_database_roundtrip.py index 25c3397ab..d8b214a69 100644 --- a/tests/test_database_roundtrip.py +++ b/tests/test_database_roundtrip.py @@ -1,30 +1,16 @@ -"""Database export/import round-trip tests (sqlite & functional_sqlite).""" +"""AB metadata loading after Excel database import (bw2io round-trip covered upstream).""" from __future__ import annotations import tempfile import time -from pathlib import Path -import pytest from bw2data.tests import bw2test from bw2io import create_core_migrations, create_default_biosphere3 -from activity_browser.bwutils.exporters import database_has_parameters from activity_browser.bwutils.metadata.loader import MDSLoader from activity_browser.bwutils.metadata.metadata import MetaDataStore -from fixtures.database_roundtrip import ( - compare_databases_semantically, - roundtrip_import, - visible_product_count, - write_source_db, -) - -BACKENDS = [ - ("sqlite", "sqlite"), - ("functional_sqlite", "functional"), -] -FORMATS = ["bw2package", "excel"] +from fixtures.database_roundtrip import roundtrip_import, visible_product_count, write_source_db def project_setup() -> None: @@ -32,35 +18,6 @@ def project_setup() -> None: create_default_biosphere3() -@bw2test -@pytest.mark.parametrize("fmt", FORMATS) -@pytest.mark.parametrize("backend,kind", BACKENDS) -@pytest.mark.parametrize("parameters", [False, True], ids=["no_params", "with_params"]) -def test_database_roundtrip(backend, kind, fmt, parameters): - project_setup() - source = f"roundtrip_{kind}_{fmt}_{'params' if parameters else 'noparams'}" - - import bw2data as bd - - write_source_db(source, kind, parameters=parameters) - assert bd.databases[source].get("backend") == backend - if parameters: - assert database_has_parameters(source) - - with tempfile.TemporaryDirectory() as tmp: - target = roundtrip_import(source, fmt, tmp) - - assert compare_databases_semantically(source, target) == [] - if kind == "functional": - assert visible_product_count(target) == 2 - - if parameters: - if fmt == "excel": - assert database_has_parameters(target) - else: - assert not database_has_parameters(target) - - @bw2test def test_load_database_populates_metadata_for_excel_import(qapp, monkeypatch): monkeypatch.setattr( diff --git a/tests/test_gsa.py b/tests/test_gsa.py index 8f5c147e8..93ea5107a 100644 --- a/tests/test_gsa.py +++ b/tests/test_gsa.py @@ -111,11 +111,6 @@ def test_get_cf_dataframe_uses_method_uncertainty(mc_project): assert "Minimum:" in dfcf.iloc[0]["uncertainty"] -def test_mc_populates_cf_dict(mc_project): - mc = _run_mc(mc_project, technosphere=False, biosphere=False, cf=True, parameters=False) - assert len(mc.CF_dict[mc.methods[0]]) == ITERATIONS - - def test_gsa_full_run_all_uncertainty_layers(mc_project_with_parameters): """End-to-end GSA with technosphere, biosphere, CF, and parameter MC uncertainty.""" mc = _run_mc(mc_project_with_parameters, **ALL_UNCERTAINTY_LAYERS) diff --git a/tests/test_lcia_overview.py b/tests/test_lcia_overview.py index 697dc7b4d..ed68c74ea 100644 --- a/tests/test_lcia_overview.py +++ b/tests/test_lcia_overview.py @@ -140,13 +140,6 @@ def test_build_flows_x_methods_flip_is_transpose_of_column_normalization(stub_me [ (0, (0, 0)), (1, (1, 1)), - (2, (1, 2)), - (3, (2, 2)), - (4, (2, 2)), - (5, (2, 3)), - (6, (2, 3)), - (8, (3, 3)), - (9, (3, 3)), (10, (3, 4)), ], ) diff --git a/tests/test_monte_carlo_uncertainty.py b/tests/test_monte_carlo_uncertainty.py index e78e1f1ec..07207d7ec 100644 --- a/tests/test_monte_carlo_uncertainty.py +++ b/tests/test_monte_carlo_uncertainty.py @@ -1,14 +1,12 @@ """ -Integration tests for ``MonteCarloLCA`` with technosphere, biosphere, CF, and parameter uncertainty. +Integration tests for ``MonteCarloLCA`` multi-layer wiring (single-layer spread is stats_arrays/bw2calc). """ from __future__ import annotations import numpy as np -import pytest from activity_browser.bwutils.montecarlo import MonteCarloLCA -from fixtures.monte_carlo import BASELINE_SCORE SEED = 42 ITERATIONS = 20 @@ -24,28 +22,6 @@ def _mc_scores(mc: MonteCarloLCA) -> np.ndarray: return mc.results[:, 0, 0] -ACTIVE_UNCERTAINTY_CASES = [ - pytest.param( - dict(technosphere=True, biosphere=False, cf=False, parameters=False), - id="technosphere", - ), - pytest.param( - dict(technosphere=False, biosphere=True, cf=False, parameters=False), - id="biosphere", - ), - pytest.param( - dict(technosphere=False, biosphere=False, cf=True, parameters=False), - id="characterization_factor", - ), -] - - -@pytest.mark.parametrize("includes", ACTIVE_UNCERTAINTY_CASES) -def test_mc_single_source_produces_spread(mc_project, includes): - mc = _run_mc(mc_project, **includes) - assert np.std(_mc_scores(mc)) > 0 - - def test_mc_parameter_uncertainty_produces_spread(mc_project_with_parameters): mc = _run_mc( mc_project_with_parameters, @@ -68,18 +44,6 @@ def test_mc_all_sources_jointly_produce_spread(mc_project_with_parameters): assert np.std(_mc_scores(mc)) > 0 -def test_mc_all_uncertainty_off_is_deterministic(mc_project): - mc = _run_mc( - mc_project, - technosphere=False, - biosphere=False, - cf=False, - parameters=False, - ) - scores = mc.results[:, 0, 0] - assert np.allclose(scores, BASELINE_SCORE) - - def test_mc_reproducible_with_same_seed(mc_project): kwargs = dict( technosphere=True, diff --git a/tests/test_project_migrate25.py b/tests/test_project_migrate25.py deleted file mode 100644 index d7c84361c..000000000 --- a/tests/test_project_migrate25.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Tests for Brightway25 project migration.""" - -import bw2data as bd - -from activity_browser.app.actions.project.project_migrate25 import MigrateThread -from fixtures.bw_helpers import write_method - - -def test_pre_process_methods_converts_legacy_tuple_cfs_to_activity_ids(basic_database): - """Legacy CF keys (database, code) must be rewritten with Brightway activity ids.""" - elementary = basic_database.get("elementary") - write_method( - "legacy_method", - [(elementary.key, 1.6)], - process=False, - ) - - MigrateThread.pre_process_methods() - - loaded = list(bd.Method(("legacy_method",)).load()) - assert loaded == [(elementary.id, 1.6)] - assert loaded[0][0] == bd.get_node(key=elementary.key).id - - -def test_pre_process_methods_preserves_regionalized_cfs(basic_database): - elementary = basic_database.get("elementary") - write_method( - "regional_method", - [(elementary.key, 1.6, "GLO")], - process=False, - ) - - MigrateThread.pre_process_methods() - - loaded = list(bd.Method(("regional_method",)).load()) - assert loaded == [(elementary.id, 1.6, "GLO")] diff --git a/tests/test_uncertainty_preview.py b/tests/test_uncertainty_preview.py new file mode 100644 index 000000000..a53a4e4e9 --- /dev/null +++ b/tests/test_uncertainty_preview.py @@ -0,0 +1,62 @@ +"""Regression tests for AB uncertainty preview helpers (not stats_arrays itself).""" +import numpy as np +import stats_arrays as sa + +from activity_browser.bwutils.uncertainty import ( + standard_uncertainty_fields, + uncertainty_reference_value, + uncertainty_statistics_scalar, +) +from activity_browser.ui.dialogs.uncertainty_pdf_preview import preview_density + + +def _lognormal_array(loc=0.0, scale=0.5, negative=False): + info = { + "uncertainty type": sa.LognormalUncertainty.id, + "loc": loc, + "scale": scale, + "shape": np.nan, + "minimum": np.nan, + "maximum": np.nan, + "negative": negative, + } + arr = sa.LognormalUncertainty.from_dicts(info) + sa.LognormalUncertainty.validate(arr) + return arr + + +def test_lognormal_preview_reference_avoids_statistics_bug(): + arr = _lognormal_array(loc=0.0, scale=0.5) + ref = uncertainty_reference_value(sa.LognormalUncertainty, arr) + assert ref == 1.0 + curve = preview_density(sa.LognormalUncertainty, arr) + assert curve is not None + assert curve.y.size > 0 + assert np.nanmax(curve.y) > 0 + + +def test_beta_pert_fields_and_preview(): + info = { + "uncertainty type": sa.BetaPERTUncertainty.id, + "minimum": 1.0, + "loc": 5.0, + "maximum": 10.0, + "scale": np.nan, + "shape": np.nan, + "negative": False, + } + assert standard_uncertainty_fields(sa.BetaPERTUncertainty.id) == [ + "minimum", + "loc", + "maximum", + "scale", + ] + arr = sa.BetaPERTUncertainty.from_dicts(info) + sa.BetaPERTUncertainty.validate(arr) + mean = uncertainty_statistics_scalar(sa.BetaPERTUncertainty, arr, "mean") + assert abs(mean - 5.166666666666667) < 1e-9 + ref = uncertainty_reference_value(sa.BetaPERTUncertainty, arr) + assert abs(ref - mean) < 1e-9 + curve = preview_density(sa.BetaPERTUncertainty, arr) + assert curve is not None + assert np.nanmax(curve.y) > 0