diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 04ef0149..9e97e00c 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -26,7 +26,7 @@ repos: exclude: ^src/badger/tests/ - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.16 + rev: v0.16.8 hooks: - id: ruff-check args: [--fix] diff --git a/AGENTS.md b/AGENTS.md index a436d2da..25b173ca 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -37,7 +37,7 @@ All tests should pass. GUI tests use `pytest-qt` and require a display server (t - **`suppress_popups` is autouse.** All tests automatically mock `ExpandableMessageBox.exec_` to prevent Qt dialogs from blocking test execution. -- **Disabled tests use `x-` prefix.** Files like `x-test_db.py` are excluded from pytest collection (they don't match the `test_*.py` pattern). These require `BADGER_DB_ROOT` to be configured. +- **Disabled tests use `x_` prefix.** Files like `x_test_db.py` are excluded from pytest collection (they don't match the `test_*.py` pattern). These require `BADGER_DB_ROOT` to be configured. - **Coverage targets the `badger` module.** `pyproject.toml` uses `--cov=badger` (module name, not a path), which resolves correctly under the `src/` layout. @@ -168,6 +168,48 @@ A pre-commit hook (`check-module-docstrings`) enforces presence. Empty `__init__ 3. **`Interface.reset_interface()`** is called after process fork — use it to reset any non-fork-safe state (file descriptors, connections, etc.) in custom interfaces. -4. **The `db.py` module is semi-deprecated** — it requires `BADGER_DB_ROOT` config which is not in the default `BadgerConfig` model. The `x-test_db.py` and `x-test_routine_id.py` files test this functionality but are excluded from normal test runs. +4. **The `db.py` module is semi-deprecated** — it requires `BADGER_DB_ROOT` config which is not in the default `BadgerConfig` model. The `x_test_db.py` and `x_test_routine_id.py` files test this functionality but are excluded from normal test runs. 5. **`utils.py` has Qt dependencies** — `BlockSignalsContext` and related utilities import from `PyQt5.QtWidgets` at the module level, so `badger.utils` cannot be imported without PyQt5 installed. + +## Conforming to Ruff linting rules + +Ruff adds many rules to the pre-commit python linting. Any of these rules can be ignored on a case-by-case basis going forward if given a valid reasoning behind the inclusion using the following format: `# noqa: - ` + +Common Ruff errors going forward. + +### BLE001 - blind-except + +```python +try: + foo() +except Exception: + ... +``` + +Use instead: + +```python +try: + foo() +except FileNotFoundError: # specific expected error to be thrown + ... +``` + +There are valid reasons to catch just the base exception in certain cases; however, catching the specific expected error will be preferred, and catching the base exception should now always include the reason for doing so. + +```python +try: + foo() +except Exception: # noqa: BLE001 - Reason for exception + ... +``` + +Alternatively, re-raising the error or exceptions logged with `exc_info` will not be flagged. + +```python +try: + foo() +except BaseException: + logger.exception("Something went wrong") +``` diff --git a/__init__.py b/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/documentation/docs/getting-started/tutorial_0.md b/documentation/docs/getting-started/tutorial_0.md index 144e484a..b36b27f2 100644 --- a/documentation/docs/getting-started/tutorial_0.md +++ b/documentation/docs/getting-started/tutorial_0.md @@ -77,24 +77,23 @@ from badger import environment class Environment(environment.Environment): - - name = 'sphere_3d' # name of the environment + name = "sphere_3d" # name of the environment variables = { # variables and their hard-limited ranges - 'x0': [-1, 1], - 'x1': [-1, 1], - 'x2': [-1, 1], + "x0": [-1, 1], + "x1": [-1, 1], + "x2": [-1, 1], } - observables = ['f'] # measurements + observables = ["f"] # measurements # Internal variables to store the current values of # the variables and observables _variables = { - 'x0': 0.0, - 'x1': 0.0, - 'x2': 0.0, + "x0": 0.0, + "x1": 0.0, + "x2": 0.0, } _observations = { - 'f': None, + "f": None, } # Variable getter -- tells Badger how to get current values of the variables @@ -109,10 +108,13 @@ class Environment(environment.Environment): self._variables[var] = x # Filling up the observations - f = self._variables['x0'] ** 2 + self._variables['x1'] ** 2 + \ - self._variables['x2'] ** 2 + f = ( + self._variables["x0"] ** 2 + + self._variables["x1"] ** 2 + + self._variables["x2"] ** 2 + ) - self._observations['f'] = [f] + self._observations["f"] = [f] # Observable getter -- how to get current values of the observables def get_observables(self, observable_names): diff --git a/documentation/docs/guides/more/create-environments-and-interfaces.md b/documentation/docs/guides/more/create-environments-and-interfaces.md index 53261b75..839ba8a1 100644 --- a/documentation/docs/guides/more/create-environments-and-interfaces.md +++ b/documentation/docs/guides/more/create-environments-and-interfaces.md @@ -41,8 +41,7 @@ from badger import interface class Interface(interface.Interface): - - name = 'myintf' + name = "myintf" def get_values(self, channel_names: list): pass @@ -98,8 +97,7 @@ from badger.interface import Interface class Environment(environment.Environment): - - name = 'myenv' + name = "myenv" variables = {} observables = [] @@ -136,32 +134,34 @@ Try to avoid doing time-consuming thing in `__init__` method. Badger would creat Okay, now we can start to implement the methods. Assume that our sample environment has 3 variables: `x`, `y`, and `z`, with range of [0, 1]. It also has 2 observations: `norm`, and `mean`. Then the `variables` and `observables` class variables should look like: ```python - variables = { - 'x': [0, 1], - 'y': [0, 1], - 'z': [0, 1], - } - observables = ['norm', 'mean'] +variables = { + "x": [0, 1], + "y": [0, 1], + "z": [0, 1], +} +observables = ["norm", "mean"] ``` Our custom env is so simple that we don't really need an interface here. Let's implement the getter and setter for the variables: ```python - # Internal variables start with a single underscore - _variables = { - 'x': 0, - 'y': 0, - 'z': 0, - } +# Internal variables start with a single underscore +_variables = { + "x": 0, + "y": 0, + "z": 0, +} - def get_variables(self, variable_names: list[str]) -> dict: - variable_outputs = {v: self._variables[v] for v in variable_names} - return variable_outputs +def get_variables(self, variable_names: list[str]) -> dict: + variable_outputs = {v: self._variables[v] for v in variable_names} - def set_variables(self, variable_inputs: dict[str, float]): - for var, x in variable_inputs.items(): - self._variables[var] = x + return variable_outputs + + +def set_variables(self, variable_inputs: dict[str, float]): + for var, x in variable_inputs.items(): + self._variables[var] = x ``` Here we use a dictionary called `_variables` to hold the values for the variables. @@ -169,19 +169,19 @@ Here we use a dictionary called `_variables` to hold the values for the variable Now let's add observable related logic: ```python - def get_observables(self, observable_names: list[str]) -> dict: - x = self._variables['x'] - y = self._variables['y'] - z = self._variables['z'] - - observable_outputs = {} - for obs in observable_names: - if obs == 'norm': - observable_outputs[obs] = (x ** 2 + y ** 2 + z ** 2) ** 0.5 - elif obs == 'mean': - observable_outputs[obs] = (x + y + z) / 3 - - return observable_outputs +def get_observables(self, observable_names: list[str]) -> dict: + x = self._variables["x"] + y = self._variables["y"] + z = self._variables["z"] + + observable_outputs = {} + for obs in observable_names: + if obs == "norm": + observable_outputs[obs] = (x**2 + y**2 + z**2) ** 0.5 + elif obs == "mean": + observable_outputs[obs] = (x + y + z) / 3 + + return observable_outputs ``` At this point, the content of `__init__.py` should be: @@ -192,21 +192,20 @@ from badger import environment class Environment(environment.Environment): - - name = 'myenv' + name = "myenv" variables = { - 'x': [0, 1], - 'y': [0, 1], - 'z': [0, 1], + "x": [0, 1], + "y": [0, 1], + "z": [0, 1], } - observables = ['norm', 'mean'] + observables = ["norm", "mean"] # Internal variables start with a single underscore _variables = { - 'x': 0, - 'y': 0, - 'z': 0, + "x": 0, + "y": 0, + "z": 0, } def get_variables(self, variable_names: list[str]) -> dict: @@ -219,15 +218,15 @@ class Environment(environment.Environment): self._variables[var] = x def get_observables(self, observable_names: list[str]) -> dict: - x = self._variables['x'] - y = self._variables['y'] - z = self._variables['z'] + x = self._variables["x"] + y = self._variables["y"] + z = self._variables["z"] observable_outputs = {} for obs in observable_names: - if obs == 'norm': - observable_outputs[obs] = (x ** 2 + y ** 2 + z ** 2) ** 0.5 - elif obs == 'mean': + if obs == "norm": + observable_outputs[obs] = (x**2 + y**2 + z**2) ** 0.5 + elif obs == "mean": observable_outputs[obs] = (x + y + z) / 3 return observable_outputs diff --git a/pyproject.toml b/pyproject.toml index 7ac825ec..92236b73 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -56,10 +56,6 @@ where = ["src"] include = ["badger"] namespaces = false -[tool.ruff.lint] -extend-select = ["TID252"] # Defaults + check imports -ignore = ["E722"] # Until bare except blocks get fixed - [tool.pytest.ini_options] addopts = "--cov=badger" log_cli_level = "INFO" diff --git a/scripts/run_tests.py b/scripts/run_tests.py old mode 100644 new mode 100755 index c1e476df..521648bb --- a/scripts/run_tests.py +++ b/scripts/run_tests.py @@ -1,9 +1,8 @@ -#!/usr/bin/env python +#!/usr/bin/env python3 import sys import pytest - if __name__ == "__main__": # Show output results from every test function # Show the message output for skipped and expected failures @@ -18,6 +17,6 @@ args.extend(["--cov=badger", "--cov-report", "term-missing"]) args.remove("--show-cov") - print("pytest arguments: {}".format(args)) + print(f"pytest arguments: {args}") print(f"Running tests on Python {sys.version}") sys.exit(pytest.main(args)) diff --git a/src/badger/__main__.py b/src/badger/__main__.py index f1d23c88..70e47f6e 100644 --- a/src/badger/__main__.py +++ b/src/badger/__main__.py @@ -5,14 +5,14 @@ import logging from badger.actions import show_info +from badger.actions.config import config_settings from badger.actions.doctor import self_check -from badger.actions.routine import show_routine -from badger.actions.generator import show_generator from badger.actions.env import show_env +from badger.actions.generator import show_generator from badger.actions.install import plugin_install -from badger.actions.uninstall import plugin_remove from badger.actions.intf import show_intf -from badger.actions.config import config_settings +from badger.actions.routine import show_routine +from badger.actions.uninstall import plugin_remove from badger.log import setup_logging logger = logging.getLogger("badger") diff --git a/src/badger/actions/__init__.py b/src/badger/actions/__init__.py index 8d6e17af..c38fc53e 100644 --- a/src/badger/actions/__init__.py +++ b/src/badger/actions/__init__.py @@ -2,6 +2,7 @@ GUI or prints config details; sub-command handlers are re-exported from their own modules (doctor, routine, env, generator, etc.).""" +import argparse import os from importlib import metadata @@ -10,27 +11,25 @@ from badger.utils import yprint -def show_info(args): +def show_info(args: argparse.Namespace) -> None: config_path = None if args.config_filepath: config_path = args.config_filepath - if args.gui or args.gui_acr: - if check_n_config_paths(args.config_filepath): - from badger.gui import launch_gui + if (args.gui or args.gui_acr) and check_n_config_paths(args.config_filepath): + from badger.gui import launch_gui - launch_gui(config_path) + launch_gui(config_path) - return + return - if args.mini: - if check_n_config_paths(args.config_filepath): - from badger.gui.mini import launch_gui + if args.mini and check_n_config_paths(args.config_filepath): + from badger.gui.mini import launch_gui - launch_gui(config_path, template_filename=args.template) + launch_gui(config_path, template_filename=args.template) - return + return if not check_n_config_paths(): return diff --git a/src/badger/actions/config.py b/src/badger/actions/config.py index e2118e58..b448b87a 100644 --- a/src/badger/actions/config.py +++ b/src/badger/actions/config.py @@ -2,11 +2,11 @@ (plugin root, archive root, theme, log level, etc.) and lets them set, skip, or reset values from the terminal.""" -from badger.settings import init_settings -import os import logging +import os -from badger.utils import yprint, convert_str_to_value +from badger.settings import init_settings +from badger.utils import convert_str_to_value, yprint logger = logging.getLogger(__name__) @@ -60,7 +60,7 @@ def _config_path_var(var_name): if _res == "y": break elif (not _res) or (_res == "n"): - print("") + print() continue else: print(f"Invalid choice: {_res}") @@ -69,7 +69,7 @@ def _config_path_var(var_name): if os.path.isdir(res): _res = input(f"Your choice is {res}, proceed ([y]/n)? ") if _res == "n": - print("") + print() continue elif (not _res) or (_res == "y"): break @@ -78,7 +78,7 @@ def _config_path_var(var_name): else: _res = input(f"{res} does not exist, do you want to create it ([y]/n)? ") if _res == "n": - print("") + print() continue elif (not _res) or (_res == "y"): os.makedirs(res) @@ -122,7 +122,7 @@ def _config_core_var(var_name): if _res == "y": break elif (not _res) or (_res == "n"): - print("") + print() continue else: print(f"Invalid choice: {_res}") diff --git a/src/badger/actions/doctor.py b/src/badger/actions/doctor.py index f8af6ac2..b2b77d06 100644 --- a/src/badger/actions/doctor.py +++ b/src/badger/actions/doctor.py @@ -2,8 +2,8 @@ offers to fix missing ones interactively, and can factory-reset Badger back to its default state.""" -from badger.settings import init_settings, mock_settings from badger.actions.config import _config_path_var +from badger.settings import init_settings, mock_settings def self_check(args): @@ -91,7 +91,7 @@ def check_n_config_paths(config_filepath=None): fixed = True for pname in issue_list: try: - print("") + print() success = _config_path_var(pname) if not success: fixed = False diff --git a/src/badger/actions/env.py b/src/badger/actions/env.py index ebee6f5d..4cfcfe6a 100644 --- a/src/badger/actions/env.py +++ b/src/badger/actions/env.py @@ -1,18 +1,20 @@ """The `badger env` command. Lists available environment plugins or shows the details (variables, observations, parameters) of a specific one.""" +import argparse import logging +from badger.errors import BadgerInvalidPluginError, BadgerPluginNotFoundError from badger.utils import range_to_str, yprint logger = logging.getLogger(__name__) -def show_env(args): +def show_env(args: argparse.Namespace) -> None: try: - from badger.factory import list_env, get_env - except Exception as e: - logger.error(e) + from badger.factory import get_env, list_env + except Exception: + logger.exception("Failed to import environment plugins.") return if args.env_name is None: @@ -21,16 +23,20 @@ def show_env(args): try: _, configs = get_env(args.env_name) - except Exception as e: + except BadgerPluginNotFoundError as e: + logger.error(e) + return + except BadgerInvalidPluginError as e: logger.error(e) - try: - # The exception could carry the configs information - configs = e.configs - except: + # The exception carries the configs information + if e.configs is None: return + configs = e.configs try: configs["variables"] = range_to_str(configs["variables"]) yprint(configs) - except: - pass + except (KeyError, TypeError): + logger.exception( + "Failed to show environment details. The configs may be malformed." + ) diff --git a/src/badger/actions/generator.py b/src/badger/actions/generator.py index 902b57e5..7f7a0059 100644 --- a/src/badger/actions/generator.py +++ b/src/badger/actions/generator.py @@ -3,6 +3,8 @@ import logging +from xopt.errors import XoptError + from badger.utils import yprint logger = logging.getLogger(__name__) @@ -10,8 +12,8 @@ def show_generator(args): try: - from badger.factory import list_generators, get_generator - except Exception as e: + from badger.factory import get_generator, list_generators + except Exception as e: # noqa: BLE001 - import triggers config/plugin loading; report and exit logger.error(e) return @@ -22,11 +24,5 @@ def show_generator(args): try: configs = get_generator(args.generator_name) yprint(configs) - except Exception as e: + except XoptError as e: logger.error(e) - try: - # The exception could carry the configs information - configs = e.configs - yprint(configs) - except: - pass diff --git a/src/badger/actions/install.py b/src/badger/actions/install.py index 92dee2ed..3bd25545 100644 --- a/src/badger/actions/install.py +++ b/src/badger/actions/install.py @@ -3,157 +3,16 @@ the current plugin installation workflow.""" import logging -import requests -import tarfile -import shutil -import os -from os.path import exists -import yaml -from tqdm.auto import tqdm - -from badger.settings import init_settings +from typing_extensions import deprecated logger = logging.getLogger(__name__) +@deprecated("The `badger install` command is currently disabled.") def plugin_install(args): print( "This command is currently disabled.\n" "Please refer to the Badger documentation for plugin management.\n\n" "Badger online documentation: https://xopt-org.github.io/Badger/" ) - return - - try: - from badger.factory import BADGER_PLUGIN_ROOT - except Exception as e: - logger.error(e) - return - - # We will not make conda as a dependency of Badger - # This is just a temp solution - # Should tell the users to use the install command a conda env is needed - from conda.cli.python_api import run_command, Commands - - hist = { - "generator": "generators", - "env": "environments", - "ext": "extensions", - "intf": "interfaces", - } - - identify = { - "optimize": "generators", - "Environment": "environments", - "Extension": "extensions", - "Interface": "interfaces", - } - - if args.plugin_type is None: - print("Please specify further what you wish to install!") - return - - if args.plugin_type != "local" and args.plugin_type not in hist: - print( - f"{args.plugin_type} is an invalid option. Choose one of the following: generator, env, ext, intf, local" - ) - return - - config = init_settings() - plugins_url = config.read_value("BADGER_PLUGINS_URL") - - if args.plugin_specific is None: - if args.plugin_type == "local": - print( - "Please provide the path to the local tarball for the plugin you wish to install" - ) - return - full_word = hist[f"{args.plugin_type}"] - url = f"{plugins_url}/api/{full_word}" - r = requests.get(url) - for elem in r.json(): - if exists(f"{BADGER_PLUGIN_ROOT}/{full_word}/{elem}"): - print(elem, " (Already installed)") - else: - print(elem) - return - - plugin_path = "" - tmp_path = os.path.join(BADGER_PLUGIN_ROOT, ".tmp") - os.makedirs(tmp_path, exist_ok=True) - - if args.plugin_type == "local": - if not exists(args.plugin_specific): - print("This local tarball does not exist!") - return - tarname = os.path.basename(os.path.normpath(args.plugin_specific)) - plugin_name = tarname[:-7] - local_path = os.path.dirname(args.plugin_specific) - os.chdir(local_path) - tar = tarfile.open(f"{tarname}", "r:gz") - tar.extractall(tmp_path) - tar.close() - - histog = {} - os.chdir(f"{tmp_path}/{plugin_name}") - with open("__init__.py", "r") as file: - info = file.read() - exec(info, histog) - identifier = list(histog.keys())[-1] - full_word = identify[identifier] - plugin_path += f"{BADGER_PLUGIN_ROOT}/{full_word}/{plugin_name}" - if exists(plugin_path): - print("This plugin is already installed!") - return - shutil.move(f"{tmp_path}/{plugin_name}", f"{BADGER_PLUGIN_ROOT}/{full_word}") - shutil.rmtree(tmp_path) - - else: - full_word = hist[f"{args.plugin_type}"] - targz_path = os.path.join(tmp_path, f"{args.plugin_specific}.tar.gz") - - r_d = requests.get(f"{plugins_url}/api/url/{full_word}/{args.plugin_specific}") - download_url = r_d.text - - r = requests.get(download_url) - if r.status_code == 200: - with open(targz_path, "wb") as f: - f.write(r.content) - os.chdir(tmp_path) - tar = tarfile.open(f"{args.plugin_specific}.tar.gz", "r:gz") - plugin_path += f"{BADGER_PLUGIN_ROOT}/{full_word}/{args.plugin_specific}" - if exists(plugin_path): - print("This plugin is already installed!") - return - print( - f"Installing {args.plugin_specific} into {BADGER_PLUGIN_ROOT}/{full_word} ..." - ) - tar.extractall(f"{BADGER_PLUGIN_ROOT}/{full_word}") - tar.close() - else: - print("The server does not have this plugin!") - return - - os.chdir(plugin_path) - with open("configs.yaml", "r") as stream: - try: - configs = yaml.safe_load(stream) - except yaml.YAMLError as e: - print(e) - else: - dependencies = configs["dependencies"] - print("Installing plugin dependencies ...") - try: - dependencies.remove("badger-opt") - except ValueError: - pass - for elem in tqdm(dependencies): - stdout_str, stderr_str, return_code_int = run_command( - Commands.INSTALL, ["-y", f"{elem}"] - ) - if return_code_int != 0: - shutil.rmtree(plugin_path) - print(stderr_str) - print("All dependencies successfully installed!") - print("Plugin installation complete!") diff --git a/src/badger/actions/intf.py b/src/badger/actions/intf.py index 2db6adfe..31d3b5f0 100644 --- a/src/badger/actions/intf.py +++ b/src/badger/actions/intf.py @@ -3,6 +3,7 @@ import logging +from badger.errors import BadgerInvalidPluginError, BadgerPluginNotFoundError from badger.utils import yprint logger = logging.getLogger(__name__) @@ -10,8 +11,8 @@ def show_intf(args): try: - from badger.factory import list_intf, get_intf - except Exception as e: + from badger.factory import get_intf, list_intf + except Exception as e: # noqa: BLE001 - import triggers config/plugin loading; report and exit logger.error(e) return @@ -22,11 +23,13 @@ def show_intf(args): try: _, configs = get_intf(args.intf_name) yprint(configs) - except Exception as e: + except BadgerPluginNotFoundError as e: logger.error(e) - try: - # The exception could carry the configs information - configs = e.configs - yprint(configs) - except: - pass + return + except BadgerInvalidPluginError as e: + logger.error(e) + # The exception carries the configs information + if e.configs is None: + logger.warning("Failed to retrieve interface configs from exception") + return + yprint(e.configs) diff --git a/src/badger/actions/routine.py b/src/badger/actions/routine.py index 2ee4bd9d..7d8ad9e5 100644 --- a/src/badger/actions/routine.py +++ b/src/badger/actions/routine.py @@ -4,61 +4,15 @@ import logging -import pandas as pd -import yaml - -from badger.utils import yprint +from typing_extensions import deprecated logger = logging.getLogger(__name__) +@deprecated("The `badger routine` command is deprecated. Please use the GUI.") def show_routine(args): print( "This command is deprecated.\n" "Please use 'badger -g' to launch the Badger GUI " "and manage routines/runs." ) - return - - try: - from badger.db import load_routine, list_routine - from badger.actions.run import run_n_archive - except Exception as e: - logger.error(e) - return - - # List routines - if args.routine_id is None: - routines = list_routine()[1] - if routines: - yprint(routines) - else: - print("No routine has been saved yet") - return - - try: - routine, _ = load_routine(args.routine_id) - if routine is None: - print(f"Routine {args.routine_id} not found") - return - except Exception as e: - print(e) - return - - # Print the routine - if not args.run: - info = yaml.safe_load(routine.yaml()) - output = {} - output["name"] = info["name"] - output["environment"] = info["environment"] - output["algorithm"] = info["generator"] - output["vocs"] = info["vocs"] - output["initial_points"] = pd.DataFrame(info["initial_points"]).to_dict("list") - output["critical_constraint_names"] = info["critical_constraint_names"] - output["tags"] = info["tags"] - output["script"] = info["script"] - - yprint(output) - return - - run_n_archive(routine, args.yes, False, args.verbose) diff --git a/src/badger/actions/run.py b/src/badger/actions/run.py index 0617e3c9..b205aff5 100644 --- a/src/badger/actions/run.py +++ b/src/badger/actions/run.py @@ -10,17 +10,18 @@ import logging import os +import signal import sys import time -import signal from pandas import DataFrame +from typing_extensions import deprecated -from badger.utils import curr_ts from badger.core import run_routine as run +from badger.errors import BadgerRunTerminated from badger.routine import Routine from badger.settings import init_settings -from badger.errors import BadgerRunTerminated +from badger.utils import curr_ts logger = logging.getLogger(__name__) @@ -30,7 +31,7 @@ def run_n_archive( ): try: from badger.archive import archive_run - except Exception as e: + except Exception as e: # noqa: BLE001 - import triggers config/plugin loading; report and exit logger.error(e) return @@ -43,7 +44,7 @@ def run_n_archive( def handler(*args): if storage["paused"]: - print("") # start a new line + print() # start a new line if flush_prompt: # erase the last prompt sys.stdout.write("\033[F") raise BadgerRunTerminated @@ -90,8 +91,8 @@ def after_evaluate(data: DataFrame): routine.environment.interface.dump_recording( os.path.join(path, filename) ) - except Exception: - pass + except Exception: # noqa: BLE001 - interface dump is best-effort + logger.warning("Failed to dump interface logs") # take a break to let the outside signal to change the status time.sleep(sleep) @@ -109,7 +110,7 @@ def states_ready(states): ) except BadgerRunTerminated as e: logger.info(e) - except Exception as e: + except Exception as e: # noqa: BLE001 - CLI run boundary logger.error(e) # Save the run when at least one solution has been evaluated @@ -120,58 +121,14 @@ def states_ready(states): path = _run["path"] filename = _run["filename"][:-4] + "pickle" routine.environment.interface.stop_recording(os.path.join(path, filename)) - except Exception: - pass + except Exception: # noqa: BLE001 - interface dump is best-effort + logger.warning("Failed to dump interface logs") +@deprecated("The `badger run` command is deprecated. Please use the GUI.") def run_routine(args): print( "This command is deprecated.\n" "Please use 'badger -g' to launch the Badger GUI " "and run an optimization." ) - return - - # try: - # from ..factory import get_algo, get_env - # except Exception as e: - # logger.error(e) - # return - - # try: - # # Get env params - # _, configs_env = get_env(args.env) - - # # Get algo params - # _, configs_algo = get_algo(args.algo) - - # # Normalize the algo and env params - # params_env = load_config(args.env_params) - # params_algo = load_config(args.algo_params) - # except Exception as e: - # logger.error(e) - # return - # params_env = merge_params(configs_env['params'], params_env) - # params_algo = merge_params(configs_algo['params'], params_algo) - - # # Load routine configs - # try: - # configs_routine = load_config(args.config) - # except Exception as e: - # logger.error(e) - # return - - # # Compose the routine - # routine = { - # 'name': args.save or generate_slug(2), - # 'algo': args.algo, - # 'env': args.env, - # 'algo_params': params_algo, - # 'env_params': params_env, - # # env_vranges is an additional info for the normalization - # # Will be removed after the normalization - # 'env_vranges': config_list_to_dict(configs_env['variables']), - # 'config': configs_routine, - # } - - # run_n_archive(routine, args.yes, args.save, args.verbose) diff --git a/src/badger/actions/uninstall.py b/src/badger/actions/uninstall.py index b21e3784..a1a0b470 100644 --- a/src/badger/actions/uninstall.py +++ b/src/badger/actions/uninstall.py @@ -1,55 +1,17 @@ """The `badger uninstall` command (currently disabled). Was intended for removing plugins — see docs for the current workflow.""" -import shutil -from os.path import exists import logging +from typing_extensions import deprecated + logger = logging.getLogger(__name__) +@deprecated("The `badger uninstall` command is currently disabled.") def plugin_remove(args): print( "This command is currently disabled.\n" "Please refer to the Badger documentation for plugin management.\n\n" "Badger online documentation: https://xopt-org.github.io/Badger/" ) - return - - try: - from badger.factory import BADGER_PLUGIN_ROOT - except Exception as e: - logger.error(e) - return - - hist = { - "generator": "generators", - "env": "environments", - "ext": "extensions", - "intf": "interfaces", - } - - if args.plugin_type is None or args.plugin_specific is None: - print("Please specify further which plugin you wish to remove!") - return - try: - full_word = hist[f"{args.plugin_type}"] - except KeyError: - logger.error(f"{args.plugin_type} is not an existing plugin type") - return - - plugin_path = f"{BADGER_PLUGIN_ROOT}/{full_word}/{args.plugin_specific}" - if not exists(plugin_path): - print(f"The plugin {args.plugin_specific} already does not exist!") - return - try: - shutil.rmtree(plugin_path) - except OSError as e: - print(f"Error: {plugin_path} : {e.strerror}") - else: - print( - f"{args.plugin_specific} was removed successfully from {BADGER_PLUGIN_ROOT}/{full_word}" - ) - print( - f"NOTE: The plugin dependencies for {args.plugin_specific} have not been removed" - ) diff --git a/src/badger/archive.py b/src/badger/archive.py index e9d39e6e..7f85e842 100644 --- a/src/badger/archive.py +++ b/src/badger/archive.py @@ -7,15 +7,15 @@ isn't lost if the process crashes. """ +import logging import os import time import warnings -import logging -from badger.utils import ts_float_to_str -from badger.settings import init_settings -from badger.routine import Routine from badger.errors import BadgerConfigError +from badger.routine import Routine +from badger.settings import init_settings +from badger.utils import ts_float_to_str logger = logging.getLogger(__name__) @@ -140,11 +140,10 @@ def list_run(): def get_runs(): runs = list_run() run_list = [] - for year, months in runs.items(): - for month, days in months.items(): - for day, files in days.items(): - for run_fname in files: - run_list.append(run_fname) + for months in runs.values(): + for days in months.values(): + for files in days.values(): + run_list.extend(files) return run_list diff --git a/src/badger/built_in_plugins/environments/sphere_2d/__init__.py b/src/badger/built_in_plugins/environments/sphere_2d/__init__.py index f9ca8a3a..0bcf25cb 100644 --- a/src/badger/built_in_plugins/environments/sphere_2d/__init__.py +++ b/src/badger/built_in_plugins/environments/sphere_2d/__init__.py @@ -2,22 +2,24 @@ objectives (f = x0^2 + x1^2, g = x0 + x1). Useful for testing optimizers without needing any hardware.""" +from typing import ClassVar + from badger import environment class Environment(environment.Environment): name = "sphere_2d" - variables = { + variables: ClassVar[dict[str, list[float]]] = { "x0": [-1, 1], "x1": [-1, 1], } - observables = ["f", "g"] + observables: ClassVar[list[str]] = ["f", "g"] - _variables = { + _variables: ClassVar[dict[str, float]] = { "x0": 0.5, "x1": 0.5, } - _observations = { + _observations: ClassVar[dict[str, float]] = { "f": 0.0, "g": 0.0, } diff --git a/src/badger/core.py b/src/badger/core.py index af636424..76c47d4b 100644 --- a/src/badger/core.py +++ b/src/badger/core.py @@ -10,9 +10,10 @@ """ import time -from typing import Callable +from collections.abc import Callable -from pandas import concat, DataFrame +from pandas import DataFrame, concat +from xopt.vocs import select_best from badger.errors import BadgerRunTerminated from badger.logger import _get_default_logger @@ -20,8 +21,6 @@ from badger.routine import Routine from badger.utils import curr_ts_to_str, dump_state -from xopt.vocs import select_best - def check_run_status(active_callback): while True: @@ -78,7 +77,7 @@ def run_routine( generate_callback: Callable, evaluate_callback: Callable, states_callback: Callable, - dump_file_callback: Callable = None, + dump_file_callback: Callable | None = None, verbose: int = 2, ) -> None: """ @@ -186,6 +185,6 @@ def run_routine( combined_results = result dump_state(dump_file, routine.generator, combined_results) - except Exception as e: + except Exception: opt_logger.update(Events.OPTIMIZATION_END, solution_meta) - raise e + raise diff --git a/src/badger/core_subprocess.py b/src/badger/core_subprocess.py index 035b3c2f..c67d7d7a 100644 --- a/src/badger/core_subprocess.py +++ b/src/badger/core_subprocess.py @@ -10,39 +10,41 @@ See core.py for the simpler in-process version of the same loop. """ -from copy import deepcopy +from __future__ import annotations + import logging +import multiprocessing as mp +import os import time import traceback -from typing import Any +from copy import deepcopy from queue import Empty +from typing import Any + from pandas import DataFrame -import multiprocessing as mp -import os +from xopt.errors import FeasibilityError, XoptError +from xopt.vocs import select_best -from badger.settings import ( - init_settings, - apply_pytorch_multiprocess_tensor_sharing_setting, -) from badger.errors import ( - BadgerRunTerminated, - BadgerEnvObsError, - MEASUREMENT_ERROR_TYPE, - MEASUREMENT_ACTION_TYPE, - MEASUREMENT_ACTION_RETRY, MEASUREMENT_ACTION_ABORT, - TERMINATION_REACHED_TYPE, - TERMINATION_ACTION_TYPE, + MEASUREMENT_ACTION_RETRY, + MEASUREMENT_ACTION_TYPE, + MEASUREMENT_ERROR_TYPE, TERMINATION_ACTION_CONTINUE, TERMINATION_ACTION_END, + TERMINATION_ACTION_TYPE, + TERMINATION_REACHED_TYPE, + BadgerEnvObsError, + BadgerRunTerminated, ) +from badger.log import configure_process_logging from badger.logger import _get_default_logger from badger.logger.event import Events from badger.routine import Routine -from badger.log import configure_process_logging -from xopt.errors import FeasibilityError, XoptError -from xopt.vocs import select_best - +from badger.settings import ( + apply_pytorch_multiprocess_tensor_sharing_setting, + init_settings, +) logger = logging.getLogger(__name__) @@ -57,7 +59,7 @@ def evaluate_measurement_with_retry( while True: try: return routine.evaluate_data(point) - except Exception as e: + except Exception as e: # noqa: BLE001 - env evaluate can raise anything error_title = f"{type(e).__name__}: {e}" error_traceback = traceback.format_exc() logger.error(f"Measurement failed: {error_title}\n{error_traceback}") @@ -192,9 +194,9 @@ def run_routine_subprocess( stop_process: mp.Event, pause_process: mp.Event, wait_event: mp.Event, - config_path: str = None, - log_queue: mp.Queue = None, - dialog_action_queue: mp.Queue = None, + config_path: str | None = None, + log_queue: mp.Queue | None = None, + dialog_action_queue: mp.Queue | None = None, ) -> None: """ Run the provided routine object using Xopt. This method is run as a subproccess @@ -231,7 +233,7 @@ def run_routine_subprocess( apply_pytorch_multiprocess_tensor_sharing_setting(config_values) # Now load the archive would use the correct config - from badger.archive import load_run, archive_run + from badger.archive import archive_run, load_run logger.info("Waiting for wait_event to be set...") wait_event.wait() @@ -240,8 +242,8 @@ def run_routine_subprocess( try: args = args_queue.get(timeout=1) logger.debug(f"Received args from queue: {args}") - except Exception as e: - logger.error(f"Error in subprocess queue.get: {type(e).__name__}, {str(e)}") + except Exception as e: # noqa: BLE001 - subprocess queue read boundary + logger.error(f"Error in subprocess queue.get: {type(e).__name__}, {e!s}") # set required arguments try: @@ -256,17 +258,16 @@ def run_routine_subprocess( routine.environment.variables.update(routine.vrange_hard_limit) # Reset data if run_data option is False - if not args["run_data"]: - if routine.data is not None: - logger.info("Resetting routine data") - routine.data = routine.data.iloc[0:0] # reset the data + if not args["run_data"] and routine.data is not None: + logger.info("Resetting routine data") + routine.data = routine.data.iloc[0:0] # reset the data except Exception as e: error_title = f"{type(e).__name__}: {e}" error_traceback = traceback.format_exc() logger.error(f"Error initializing routine: {error_title}\n{error_traceback}") queue.put((error_title, error_traceback)) - raise e + raise # TODO look into this bug with serializing of turbo. Fix might be needed in Xopt # Patch for converting dtype str to torch object @@ -276,13 +277,10 @@ def run_routine_subprocess( routine.generator.turbo_controller.tkwargs["dtype"] = eval(dtype) except AttributeError: logger.warning("AttributeError when converting turbo_controller dtype") - pass except KeyError: logger.warning("KeyError when converting turbo_controller dtype") - pass except TypeError: logger.warning("TypeError when converting turbo_controller dtype") - pass # Assign the initial points and bounds logger.info(f"Setting routine variable ranges: {args['variable_ranges']}") @@ -436,10 +434,9 @@ def run_routine_subprocess( logger.debug("Sending evaluation data to evaluate_queue.") evaluate_queue[0].send((routine.data, generator_copy)) - if archive: - if not testing: - logger.info("Archiving run state.") - archive_run(routine) + if archive and not testing: + logger.info("Archiving run state.") + archive_run(routine) except BadgerRunTerminated: logger.info("Optimization terminated by BadgerRunTerminated.") @@ -461,4 +458,4 @@ def run_routine_subprocess( error_traceback = traceback.format_exc() queue.put((error_title, error_traceback)) evaluate_queue[0].close() - raise e + raise diff --git a/src/badger/db.py b/src/badger/db.py index 635ac767..28d8b6dd 100644 --- a/src/badger/db.py +++ b/src/badger/db.py @@ -6,19 +6,19 @@ in settings) and can be exported/imported for sharing between installations. """ +import logging import os +import sqlite3 +import uuid import warnings -from datetime import datetime -import logging +from datetime import UTC, datetime import yaml -import sqlite3 -import uuid +from badger.errors import BadgerConfigError, BadgerDBError from badger.routine import Routine from badger.settings import init_settings from badger.utils import get_yaml_string -from badger.errors import BadgerConfigError, BadgerDBError logger = logging.getLogger(__name__) @@ -91,8 +91,8 @@ def filter_routines(records, tags): _tags = yaml.safe_load(record[3])["config"]["tags"] if tags.items() <= _tags.items(): records_filtered.append(record) - except: - pass + except (KeyError, TypeError, yaml.YAMLError) as e: + logger.warning(f"Failed to extract tags from routine {record[0]}: {e}") return records_filtered @@ -107,7 +107,7 @@ def extract_metadata(records): env_list.append(env) descr = metadata["description"] descr_list.append(descr) - except Exception: + except (KeyError, TypeError, yaml.YAMLError): env_list.append("") descr_list.append("") @@ -125,7 +125,7 @@ def save_routine(routine: Routine): routine.id = id cur.execute( "insert into routine values (?, ?, ?, ?)", - (routine.id, routine.name, routine.yaml(), datetime.now()), + (routine.id, routine.name, routine.yaml(), datetime.now(tz=UTC)), ) con.commit() @@ -146,7 +146,7 @@ def update_routine(routine: Routine): if record: # update the record cur.execute( "update routine set name = ?, config = ?, savedAt = ? where id = ?", - (routine.name, routine.yaml(), datetime.now(), routine.id), + (routine.name, routine.yaml(), datetime.now(tz=UTC), routine.id), ) con.commit() @@ -215,26 +215,25 @@ def load_routine(id: str): @ensure_routines_db_exists -def list_routine(keyword="", tags={}): +def list_routine(keyword="", tags: dict[str, str] | None = None): + if tags is None: + tags = {} db_routine = os.path.join(BADGER_DB_ROOT, "routines.db") con = sqlite3.connect(db_routine) cur = con.cursor() - # check if id column is in database # if not, add it and update routine and run entries accordingly cur.execute("pragma table_info(routine)") columns = [row[1] for row in cur.fetchall()] if "id" not in columns: - cur.execute( - """ + cur.execute(""" create table new_table ( id text primary key, name text, config, savedAt timestamp ) - """ - ) + """) db_run = os.path.join(BADGER_DB_ROOT, "runs.db") con_run = sqlite3.connect(db_run) cur_run = con_run.cursor() @@ -318,8 +317,8 @@ def save_run(run): routine_id = run["routine"].id run_filename = run["filename"] timestamps = run["data"]["timestamp"] - time_start = datetime.fromtimestamp(timestamps[0]) - time_finish = datetime.fromtimestamp(timestamps[-1]) + time_start = datetime.fromtimestamp(timestamps[0], tz=UTC) + time_finish = datetime.fromtimestamp(timestamps[-1], tz=UTC) # Check if the record exist (same filename) cur.execute("select id from run where filename = ?", (run_filename,)) @@ -425,7 +424,7 @@ def import_routines(filename): for record in records: try: cur_db.execute("insert into routine values (?, ?, ?, ?)", record) - except: + except sqlite3.Error: failed_list.append(record[0]) con_db.commit() diff --git a/src/badger/environment.py b/src/badger/environment.py index de99649e..bfd7d628 100644 --- a/src/badger/environment.py +++ b/src/badger/environment.py @@ -10,12 +10,13 @@ evaluation on computed observables (see formula.py). """ +import logging from abc import abstractmethod -from logging import warning -from typing import TYPE_CHECKING, Any, ClassVar, Dict, List, Optional +from typing import TYPE_CHECKING, Any, ClassVar from pydantic import BaseModel, ConfigDict, Field, SerializeAsAny from pydantic._internal._model_construction import ModelMetaclass + from badger.errors import ( BadgerEnvVarError, BadgerNoInterfaceError, @@ -23,12 +24,15 @@ if TYPE_CHECKING: from badger.factory import BadgerPluginConfig + from badger.formula import extract_variable_keys, interpret_expression from badger.interface import Interface +logger = logging.getLogger(__name__) + def validate_setpoints(func): - def validate(cls, variable_inputs: Dict[str, float]): + def validate(cls, variable_inputs: dict[str, float]): _bounds = cls.get_bounds(list(variable_inputs.keys())) for name, value in variable_inputs.items(): lower = _bounds[name][0] @@ -51,7 +55,7 @@ def process_formulas(func): to process formulas if they exist in the observable names. """ - def process(cls, observable_names: List[str]) -> Dict[str, float]: + def process(cls, observable_names: list[str]) -> dict[str, float]: # get the list of observable names needed by themselves and any formulas formula_observables = [] basic_observables = [] @@ -92,7 +96,7 @@ def process(cls, observable_names: List[str]) -> Dict[str, float]: def validate_bounds(func): - def validate(cls, variable_names: List[str]): + def validate(cls, variable_names: list[str]): bounds = func(cls, variable_names) for name, bound in bounds.items(): @@ -142,11 +146,11 @@ class BaseEnvironment(BaseModel, metaclass=EnvMeta): validate_assignment=True, use_enum_values=True, arbitrary_types_allowed=True ) name: ClassVar[str] = Field(description="environment name") - variables: ClassVar[Dict[str, list[float]]] # bounds list could be empty for var + variables: ClassVar[dict[str, list[float]]] # bounds list could be empty for var observables: ClassVar[list[str]] @abstractmethod - def get_variables(self, variable_names: list[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: """ Get the values of the specified variables from the environment. @@ -160,10 +164,9 @@ def get_variables(self, variable_names: list[str]) -> Dict[str, float]: Dict[str, float] A dictionary mapping variable names to their values. """ - pass @abstractmethod - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): """ Set the values of the specified variables in the environment. @@ -172,12 +175,11 @@ def set_variables(self, variable_inputs: Dict[str, float]): variable_inputs : Dict[str, float] A dictionary mapping variable names to their values. """ - pass @abstractmethod def get_observables( - self, observable_names: List[str] - ) -> Dict[str, float | List[float]]: + self, observable_names: list[str] + ) -> dict[str, float | list[float]]: """ Get the values of the specified observables from the environment. @@ -195,16 +197,14 @@ def get_observables( A dictionary mapping observable names to their values. """ - pass def reset_environment(self): """ Reset the environment to its initial state. This method is called at the start of each run. """ - pass - def get_system_states(self) -> Dict[str, Any]: + def get_system_states(self) -> dict[str, Any]: """ Get the current system states from the environment. This method is called to retrieve the current state of the environment. @@ -299,7 +299,7 @@ def get_observable(self, observable_name: str) -> float: class Environment(BaseEnvironment): # Interface - interface: Optional[SerializeAsAny[Interface]] = None + interface: SerializeAsAny[Interface] | None = None # Put all other env params here # params: float = Field(..., description='Example env parameter') @@ -307,19 +307,19 @@ class Environment(BaseEnvironment): # Optional methods to inherit ############################################################ - def get_variables(self, variable_names: List[str]) -> Dict: + def get_variables(self, variable_names: list[str]) -> dict: if not self.interface: raise BadgerNoInterfaceError return self.interface.get_values(variable_names) - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): if not self.interface: raise BadgerNoInterfaceError return self.interface.set_values(variable_inputs) - def get_observables(self, observable_names: List[str]) -> Dict: + def get_observables(self, observable_names: list[str]) -> dict: if not self.interface: raise BadgerNoInterfaceError @@ -329,7 +329,7 @@ def reset_environment(self): if self.interface: return self.interface.reset_interface() - def get_info(self, variable_names: List[str]) -> Dict | None: + def get_info(self, variable_names: list[str]) -> dict | None: if not self.interface: return None @@ -353,8 +353,8 @@ def instantiate_env( intf_name = configs["interface"][0] except KeyError: intf_name = None - except Exception as e: - warning(e) + except (TypeError, IndexError) as e: + logger.warning(e) intf_name = None if intf_name is not None: diff --git a/src/badger/errors.py b/src/badger/errors.py index 02a575d4..c7a8efed 100644 --- a/src/badger/errors.py +++ b/src/badger/errors.py @@ -2,9 +2,11 @@ with the traceback when raised inside the GUI. Subclasses cover config issues, database errors, plugin failures, and optimization stop signals.""" -from PyQt5.QtWidgets import QMessageBox -import traceback import sys +import traceback +from typing import Any + +from PyQt5.QtWidgets import QMessageBox class BadgerError(Exception): @@ -86,7 +88,9 @@ class BadgerInterfaceChannelError(Exception): class BadgerInvalidPluginError(Exception): - pass + def __init__(self, message: str = "", configs: Any = None): + super().__init__(message) + self.configs = configs class BadgerPluginNotFoundError(Exception): diff --git a/src/badger/extension.py b/src/badger/extension.py index 3ef03169..1d60e312 100644 --- a/src/badger/extension.py +++ b/src/badger/extension.py @@ -3,7 +3,6 @@ Third-party packages can subclass Extension to plug in custom algorithms.""" from abc import ABC, abstractmethod -from typing import List class Extension(ABC): @@ -18,7 +17,7 @@ def __init__(self): # List all available generators @abstractmethod - def list_generator(self) -> List[str]: + def list_generator(self) -> list[str]: pass # Get config of an generator diff --git a/src/badger/factory.py b/src/badger/factory.py index 5c0f0e1e..cfa2e5fe 100644 --- a/src/badger/factory.py +++ b/src/badger/factory.py @@ -96,7 +96,7 @@ def scan_plugins(root: str): for fname in os.listdir(proot) if os.path.exists(os.path.join(proot, fname, "__init__.py")) ] - except: + except OSError: plugins = [] for pname in plugins: @@ -136,11 +136,10 @@ def load_plugin( try: module = importlib.import_module(f"{ptype}s.{pname}") except ImportError as e: - _e = BadgerInvalidPluginError( - f"{ptype} {pname} is not available due to missing dependencies: {e}" - ) - _e.configs = configs # attach information to the exception - raise _e + raise BadgerInvalidPluginError( + f"{ptype} {pname} is not available due to missing dependencies: {e}", + configs=configs, + ) from e if ptype == "generator": plugin = (module.optimize, configs) @@ -172,7 +171,7 @@ def load_plugin( intf = cast(BadgerInterface, Interface()) except KeyError: intf = None - except Exception as e: + except Exception as e: # noqa: BLE001 - interface plugin load best-effort logger.warning(e) intf = None env = m_env(interface=intf, params=configs) @@ -196,7 +195,7 @@ def load_plugin( return plugin -def load_badger_docs(name: str, ptype: str = None) -> str: +def load_badger_docs(name: str, ptype: str | None = None) -> str: """ Load general Badger documentation from Badger/documentation/docs/guides. @@ -204,8 +203,8 @@ def load_badger_docs(name: str, ptype: str = None) -> str: __________ name : str Name of the .md file to open - subdir : str (None) - Name of subdirectory if file is not in main guides directory + ptype : str | None (None) + Type of plugin (e.g., 'generator', 'interface', 'environment') Returns _______ @@ -236,7 +235,7 @@ def load_badger_docs(name: str, ptype: str = None) -> str: try: with open(docs_dir / f"{name}.md", "r") as f: readme = f.read() - except: + except OSError: readme = f"# {name}\nNo documentation found.\n" if ptype == "generator": @@ -295,10 +294,10 @@ def load_plugin_docs(pname: str, ptype: str) -> str: docstring = module.Environment.__doc__ return _format_docs_str(readme, docstring, ptype) - except: + except Exception as e: raise BadgerInvalidDocsError( f"Error loading docs for {ptype} {pname}: docs not found" - ) + ) from e def _format_docs_str(readme: str, docstring: str, ptype: str) -> str: @@ -356,7 +355,7 @@ def _format_md_docs(text: str): def _md_images_to_html( text: str, - base_prefix: str = None, + base_prefix: str | None = None, width: int = 575, ) -> str: """ diff --git a/src/badger/formula.py b/src/badger/formula.py index 5be6336b..51208bbd 100644 --- a/src/badger/formula.py +++ b/src/badger/formula.py @@ -7,10 +7,11 @@ namespace (numpy only). Includes typo detection for misspelled variable names. """ -import numpy as np -import re import ast import difflib +import re + +import numpy as np def safe_var_name(var_name): @@ -97,4 +98,4 @@ def interpret_expression(expr, variables): try: return eval(expr, {"__builtins__": {}}, safe_namespace) except Exception as e: - raise ValueError(f"Expression evaluation failed: {e}") + raise ValueError(f"Expression evaluation failed: {e}") from e diff --git a/src/badger/gui/__init__.py b/src/badger/gui/__init__.py index 45976b40..7803fc8e 100644 --- a/src/badger/gui/__init__.py +++ b/src/badger/gui/__init__.py @@ -1,23 +1,23 @@ """Launches the Badger GUI. Sets up the QApplication (theming, DPI scaling, error handling) and opens the main window. Called via `badger -g`.""" -from importlib import resources +import logging import signal import sys import time -from PyQt5.QtWidgets import QApplication -from PyQt5.QtGui import QFont, QIcon -from PyQt5 import QtCore -from qdarkstyle import load_stylesheet, LightPalette, DarkPalette +import traceback +from importlib import resources +from types import TracebackType +from typing import NoReturn -from badger.settings import init_settings -from badger.gui.windows.main_window import BadgerMainWindow +from PyQt5 import QtCore +from PyQt5.QtGui import QFont, QIcon +from PyQt5.QtWidgets import QApplication +from qdarkstyle import DarkPalette, LightPalette, load_stylesheet -import traceback from badger.errors import BadgerError -from types import TracebackType -from typing import Type, NoReturn -import logging +from badger.gui.windows.main_window import BadgerMainWindow +from badger.settings import init_settings logger = logging.getLogger(__name__) @@ -60,7 +60,7 @@ def on_timeout(): def error_handler( - etype: Type[BaseException], value: BaseException, tb: TracebackType + etype: type[BaseException], value: BaseException, tb: TracebackType ) -> NoReturn: """ Custom exception handler that formats uncaught exceptions and raises a BadgerError. diff --git a/src/badger/gui/components/action_bar.py b/src/badger/gui/components/action_bar.py index e39d827f..7bb58b9f 100644 --- a/src/badger/gui/components/action_bar.py +++ b/src/badger/gui/components/action_bar.py @@ -1,11 +1,20 @@ """Toolbar with run-control buttons (start, pause, stop), logbook submission, docs access, and the extensions palette launcher.""" -from PyQt5.QtWidgets import QStyle, QStyleOptionToolButton, QWidget, QHBoxLayout -from PyQt5.QtWidgets import QToolButton, QMenu, QAction -from PyQt5.QtGui import QIcon, QFont -from PyQt5.QtCore import QEvent, pyqtSignal, QSize from importlib import resources + +from PyQt5.QtCore import QEvent, QSize, pyqtSignal +from PyQt5.QtGui import QFont, QIcon +from PyQt5.QtWidgets import ( + QAction, + QHBoxLayout, + QMenu, + QStyle, + QStyleOptionToolButton, + QToolButton, + QWidget, +) + from badger.gui.utils import create_button from badger.gui.windows.docs_window import BadgerDocsWindow diff --git a/src/badger/gui/components/analysis_extensions.py b/src/badger/gui/components/analysis_extensions.py index 4bf6b1d5..8246a19e 100644 --- a/src/badger/gui/components/analysis_extensions.py +++ b/src/badger/gui/components/analysis_extensions.py @@ -2,7 +2,7 @@ viewer). Each dialog receives live data updates from the run monitor.""" import logging -from typing import Optional, cast +from typing import cast from PyQt5.QtCore import Qt, pyqtSignal from PyQt5.QtGui import QCloseEvent @@ -26,7 +26,7 @@ class AnalysisExtension(QWidget): generator_type: type[Generator] widget: AnalysisWidget - def __init__(self, parent: Optional[QWidget] = None): + def __init__(self, parent: QWidget | None = None): super().__init__(parent=parent) # A parented QWidget is a child widget by default. Setting the Window # flag keeps the palette as the owner (for lifetime/stacking) while @@ -86,7 +86,7 @@ class ParetoFrontViewer(AnalysisExtension): def __init__( self, routine: Routine, - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): super().__init__(parent=parent) @@ -101,7 +101,7 @@ class BOVisualizer(AnalysisExtension): def __init__( self, routine: Routine, - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): super().__init__(parent=parent) @@ -118,7 +118,7 @@ class BaxVisualizer(AnalysisExtension): def __init__( self, routine: Routine, - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): super().__init__(parent=parent) diff --git a/src/badger/gui/components/analysis_widget.py b/src/badger/gui/components/analysis_widget.py index ee4c3b56..d300e7d2 100644 --- a/src/badger/gui/components/analysis_widget.py +++ b/src/badger/gui/components/analysis_widget.py @@ -4,7 +4,7 @@ import logging from abc import abstractmethod from collections.abc import Callable -from typing import Any, Optional +from typing import Any from PyQt5.QtWidgets import QWidget from xopt import Generator @@ -19,7 +19,7 @@ class AnalysisWidget(QWidget): routine: Routine generator: Generator - parameters: dict[str, Any] = {} + df_length: float = float("inf") initialized: bool = False routine_identifier: str = "" @@ -30,8 +30,9 @@ class AnalysisWidget(QWidget): def __init__( self, routine: Routine, - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): + self.parameters: dict[str, Any] = {} super().__init__(parent=parent) self.routine = routine self.generator = routine.generator diff --git a/src/badger/gui/components/archive_search.py b/src/badger/gui/components/archive_search.py index 541eb95a..1449e0a3 100644 --- a/src/badger/gui/components/archive_search.py +++ b/src/badger/gui/components/archive_search.py @@ -3,8 +3,7 @@ when building a routine.""" import logging -from typing import List -from PyQt5.QtGui import QDrag, QKeyEvent + from PyQt5.QtCore import ( QAbstractTableModel, QMimeData, @@ -14,6 +13,7 @@ QVariant, pyqtSignal, ) +from PyQt5.QtGui import QDrag, QKeyEvent from PyQt5.QtWidgets import ( QAbstractItemView, QHBoxLayout, @@ -25,6 +25,7 @@ QVBoxLayout, QWidget, ) + from badger.errors import BadgerRoutineError logger = logging.getLogger(__name__) @@ -89,7 +90,7 @@ def append(self, pv: str) -> None: self.endInsertRows() self.layoutChanged.emit() - def replace_rows(self, pvs: List[str]) -> None: + def replace_rows(self, pvs: list[str]) -> None: """Overwrites any existing rows in the table with the input list of variable names""" self.beginInsertRows(QModelIndex(), 0, len(pvs) - 1) self.results_list = pvs diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index 37dee4d5..a1325eaf 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -10,7 +10,6 @@ import logging import time from dataclasses import dataclass, field -from typing import Optional from PyQt5.QtWidgets import QSizePolicy, QVBoxLayout, QWidget from xopt.generators.bayesian.bax_generator import BaxGenerator @@ -74,7 +73,7 @@ class BaxWidget(AnalysisWidget): generator: BaxGenerator parameters: Parameters # type: ignore[assignment] - def __init__(self, routine: Routine, parent: Optional[QWidget] = None): + def __init__(self, routine: Routine, parent: QWidget | None = None): logger.debug("Initializing BaxWidget") super().__init__(routine=routine, parent=parent) diff --git a/src/badger/gui/components/bax_visualizer/controls.py b/src/badger/gui/components/bax_visualizer/controls.py index ec22a412..fee4bee5 100644 --- a/src/badger/gui/components/bax_visualizer/controls.py +++ b/src/badger/gui/components/bax_visualizer/controls.py @@ -4,7 +4,7 @@ for variable selection and visualization updates in the BAX visualizer. """ -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( @@ -34,14 +34,13 @@ class ControlsWidget(QWidget): - ref_inputs: list[QTableWidgetItem] = [] - def __init__( self, routine: Routine, parameters: "Parameters", - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ) -> None: + self.ref_inputs: list[QTableWidgetItem] = [] super().__init__(parent=parent) self.routine = routine self.parameters = parameters diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index c5a5fc76..72504053 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -1,9 +1,14 @@ """Matplotlib-based plotting widget for visualizing BAX virtual measurements.""" import logging -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING import matplotlib.pyplot as plt +from bax_algorithms.visualize import ( + plot_bax_input_convergence, + plot_bax_objective_convergence, + visualize_virtual_measurement_result, +) from matplotlib.axes import Axes from matplotlib.backends.backend_qt import NavigationToolbar2QT from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg @@ -21,11 +26,6 @@ clear_tabs, ) from badger.utils import BlockSignalsContext -from bax_algorithms.visualize import ( # noqa: E402 - plot_bax_input_convergence, - plot_bax_objective_convergence, - visualize_virtual_measurement_result, -) if TYPE_CHECKING: from badger.gui.components.bax_visualizer.bax_widget import Parameters @@ -39,7 +39,7 @@ def __init__( self, generator: BaxGenerator, parameters: "Parameters", - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): logger.debug("Initializing PlottingWidget") super().__init__(parent=parent) @@ -177,7 +177,7 @@ def update_first_tab(self) -> None: content = self._build_plot_widget(fig, add_stretch=True) self.plot_tab_widget.addTab(self._scrollable(content), "Virtual Objective") plt.close(fig) - except Exception as e: + except Exception as e: # noqa: BLE001 - plot render can fail many ways logger.error(f"Error creating plot: {e}") self.plot_tab_widget.addTab(QWidget(), "Error") @@ -190,7 +190,7 @@ def update_second_tab(self) -> None: fig, _ = create_plot() layout.addWidget(self._build_plot_widget(fig)) plt.close(fig) - except Exception as e: + except Exception as e: # noqa: BLE001 - plot render can fail many ways logger.error(f"Error creating plot: {e}") # Keep the plots top-aligned at their natural height; the scroll area diff --git a/src/badger/gui/components/bax_visualizer/ui.py b/src/badger/gui/components/bax_visualizer/ui.py index e7b1d75f..b8d02a1b 100644 --- a/src/badger/gui/components/bax_visualizer/ui.py +++ b/src/badger/gui/components/bax_visualizer/ui.py @@ -1,7 +1,6 @@ """UI layout definitions for the BAX visualizer widget.""" -from typing import TYPE_CHECKING, Optional - +from typing import TYPE_CHECKING from badger.gui.components.bax_visualizer.controls import ControlsWidget from badger.gui.components.extension_utilities import ( @@ -22,7 +21,7 @@ def __init__( self, routine: Routine, parameters: "Parameters", - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): super().__init__(parent=parent) diff --git a/src/badger/gui/components/bo_visualizer/bo_widget.py b/src/badger/gui/components/bo_visualizer/bo_widget.py index c36db762..bfba76d6 100644 --- a/src/badger/gui/components/bo_visualizer/bo_widget.py +++ b/src/badger/gui/components/bo_visualizer/bo_widget.py @@ -8,7 +8,7 @@ """ import logging -from typing import Optional, cast +from typing import cast from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( @@ -59,16 +59,18 @@ class BOPlotWidget(AnalysisWidget): generator: BayesianGenerator # pyright: ignore[reportIncompatibleVariableOverride] - parameters: ConfigurableOptions = DEFAULT_PARAMETERS.copy() + df_length: float = float("inf") initialized: bool = False def __init__( self, routine: Routine, - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): logger.debug("Initializing BOPlotWidget") + + self.parameters: ConfigurableOptions = DEFAULT_PARAMETERS.copy() super().__init__(routine, parent) self.create_ui() @@ -218,7 +220,7 @@ def on_set_best_reference_point_clicked(self) -> None: logger.debug("Setting best reference points") try: self.set_best_reference_points() - except Exception as e: + except Exception as e: # noqa: BLE001 - UI action reports via dialog logger.error(f"Error getting best reference points: {e}") QMessageBox.critical( self, @@ -231,7 +233,7 @@ def on_set_latest_reference_points_clicked(self) -> None: logger.debug("Setting latest reference points") try: self.set_latest_reference_points() - except Exception as e: + except Exception as e: # noqa: BLE001 - UI action reports via dialog logger.error(f"Error getting latest reference points: {e}") QMessageBox.critical( self, @@ -507,10 +509,10 @@ def update_routine(self, routine: Routine, generator_type: type[Generator]) -> N self.generator.train_model(self.routine.data) except HandledException as he: logger.error(str(he)) - raise he + raise except Exception as e: logger.error(str(e)) - raise e + raise def set_best_reference_points( self, diff --git a/src/badger/gui/components/bo_visualizer/plotting_area.py b/src/badger/gui/components/bo_visualizer/plotting_area.py index 47a50e3b..b593669d 100644 --- a/src/badger/gui/components/bo_visualizer/plotting_area.py +++ b/src/badger/gui/components/bo_visualizer/plotting_area.py @@ -4,7 +4,7 @@ import logging import time from collections.abc import Callable -from typing import Optional, cast +from typing import cast from matplotlib.axes import Axes from matplotlib.backends.backend_qt import NavigationToolbar2QT as NavigationToolbar @@ -33,9 +33,9 @@ class PlottingArea(QWidget): - last_updated: Optional[float] = None + last_updated: float | None = None - def __init__(self, parent: Optional[QWidget] = None): + def __init__(self, parent: QWidget | None = None): super().__init__(parent) # Create a layout for the plot area without pre-filling it with a plot @@ -111,10 +111,10 @@ def update_plot( # Add the new canvas to the layout layout.addWidget(canvas) layout.addWidget(toolbar) - except HandledException as he: - raise he + except HandledException: + raise except Exception as e: logger.error(f"Error updating plot: {e}") - raise e + raise # Update the last updated time self.last_updated = time.time() diff --git a/src/badger/gui/components/bo_visualizer/ui_components.py b/src/badger/gui/components/bo_visualizer/ui_components.py index e8f9ea57..97dd30eb 100644 --- a/src/badger/gui/components/bo_visualizer/ui_components.py +++ b/src/badger/gui/components/bo_visualizer/ui_components.py @@ -29,12 +29,11 @@ class UIComponents: - variables: list[str] = [] - def __init__( self, default_parameters: ConfigurableOptions, ): + self.variables: list[str] = [] self.variable_checkboxes: dict[str, QCheckBox] = {} self.ref_inputs: list[QTableWidgetItem] = [] self.reference_table = QTableWidget() diff --git a/src/badger/gui/components/bounds_preview.py b/src/badger/gui/components/bounds_preview.py index 5b1c9b09..a8610f3a 100644 --- a/src/badger/gui/components/bounds_preview.py +++ b/src/badger/gui/components/bounds_preview.py @@ -1,7 +1,7 @@ """Widget for visualizing hard bounds, preview scan bounds, and current value.""" -from PyQt5.QtCore import Qt, QRectF -from PyQt5.QtGui import QColor, QPainter, QPen, QBrush, QPaintEvent +from PyQt5.QtCore import QRectF, Qt +from PyQt5.QtGui import QBrush, QColor, QPainter, QPaintEvent, QPen from PyQt5.QtWidgets import QWidget @@ -79,22 +79,18 @@ def paintEvent(self, _event: QPaintEvent) -> None: painter.drawRoundedRect(preview_rect, 1, 1) # current line - curr_x = int(round(self._to_x(self.curr, left, track_w))) + curr_x = round(self._to_x(self.curr, left, track_w)) painter.setPen(QPen(QColor(240, 240, 240), 1)) - painter.drawLine( - curr_x, int(round(track_top)), curr_x, int(round(track_top + track_h)) - ) + painter.drawLine(curr_x, round(track_top), curr_x, round(track_top + track_h)) # text labels font = painter.font() font.setPointSizeF(10.0) painter.setFont(font) painter.setPen(QPen(QColor(190, 190, 190), 1)) - label_y = int(round(track_top + track_h + 12)) - painter.drawText(int(round(left)), label_y, f"{self.hard_lower:.2f}") - painter.drawText( - int(round(max(left, right - 28))), label_y, f"{self.hard_upper:.2f}" - ) + label_y = round(track_top + track_h + 12) + painter.drawText(round(left), label_y, f"{self.hard_lower:.2f}") painter.drawText( - int(round(max(left, curr_x - 16))), label_y, f"{self.curr:.3f}" + round(max(left, right - 28)), label_y, f"{self.hard_upper:.2f}" ) + painter.drawText(round(max(left, curr_x - 16)), label_y, f"{self.curr:.3f}") diff --git a/src/badger/gui/components/collapsible_box.py b/src/badger/gui/components/collapsible_box.py index 85897932..0f67e88e 100644 --- a/src/badger/gui/components/collapsible_box.py +++ b/src/badger/gui/components/collapsible_box.py @@ -5,7 +5,6 @@ from PyQt5 import QtCore, QtGui, QtWidgets from PyQt5.QtWidgets import QLayout - stylesheet_toolbutton = """ QToolButton { @@ -20,13 +19,13 @@ class ScrollArea(QtWidgets.QScrollArea): def resizeEvent(self, e): self.resized.emit() - return super(ScrollArea, self).resizeEvent(e) + return super().resizeEvent(e) # https://stackoverflow.com/a/52617714/4263605 class CollapsibleBox(QtWidgets.QWidget): def __init__(self, parent=None, title="", duration=100, tooltip=None): - super(CollapsibleBox, self).__init__(parent) + super().__init__(parent) self.title = title self.duration = duration diff --git a/src/badger/gui/components/con_table.py b/src/badger/gui/components/con_table.py index 10b61c3d..bd44b928 100644 --- a/src/badger/gui/components/con_table.py +++ b/src/badger/gui/components/con_table.py @@ -3,12 +3,13 @@ reordering.""" from typing import Any + from PyQt5.QtWidgets import ( QCheckBox, QComboBox, - QStyledItemDelegate, QDoubleSpinBox, QMessageBox, + QStyledItemDelegate, QWidget, ) diff --git a/src/badger/gui/components/constraint_item.py b/src/badger/gui/components/constraint_item.py index b4882dc7..8574212c 100644 --- a/src/badger/gui/components/constraint_item.py +++ b/src/badger/gui/components/constraint_item.py @@ -1,15 +1,18 @@ """Creates a single constraint row widget: observable selector, relation combo (>, <, =), threshold spinbox, criticality checkbox, and remove button.""" +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( + QAbstractSpinBox, + QCheckBox, + QComboBox, + QDoubleSpinBox, QHBoxLayout, QPushButton, + QStyledItemDelegate, QWidget, - QDoubleSpinBox, - QAbstractSpinBox, ) -from PyQt5.QtWidgets import QComboBox, QCheckBox, QStyledItemDelegate -from PyQt5.QtCore import Qt + from badger.gui.utils import ( MouseWheelWidgetAdjustmentGuard, NoHoverFocusComboBox, @@ -30,7 +33,7 @@ def constraint_item( cb_obs.addItems(options) try: idx = options.index(name) - except: + except ValueError: idx = 0 cb_obs.setCurrentIndex(idx) diff --git a/src/badger/gui/components/create_process.py b/src/badger/gui/components/create_process.py index b32b3e35..a263ce87 100644 --- a/src/badger/gui/components/create_process.py +++ b/src/badger/gui/components/create_process.py @@ -1,15 +1,14 @@ """QThread worker that pre-spawns an optimization subprocess in the background so it's ready to go when the user hits "start".""" +import logging from multiprocessing import Event, Pipe, Process, Queue -from PyQt5.QtCore import pyqtSignal, QObject +from PyQt5.QtCore import QObject, pyqtSignal -from badger.settings import init_settings from badger.core_subprocess import run_routine_subprocess - -import logging from badger.log import get_logging_manager +from badger.settings import init_settings logger = logging.getLogger(__name__) diff --git a/src/badger/gui/components/data_panel.py b/src/badger/gui/components/data_panel.py index ba7ec369..e64a339b 100644 --- a/src/badger/gui/components/data_panel.py +++ b/src/badger/gui/components/data_panel.py @@ -1,30 +1,31 @@ """Panel for viewing and managing pre-loaded optimization data. Lets users load data from archived runs or clear the buffer before starting.""" +import pandas as pd +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( - QVBoxLayout, + QCheckBox, + QGroupBox, QHBoxLayout, + QLabel, + QMessageBox, QPushButton, + QVBoxLayout, QWidget, - QMessageBox, ) -import pandas as pd -from PyQt5.QtWidgets import QGroupBox, QCheckBox, QLabel -from PyQt5.QtCore import Qt +from xopt.vocs import VOCS + from badger.gui.components.data_table import ( TableWithCopy, -) -from badger.gui.components.data_table import ( data_table, get_horizontal_header_as_list, - update_table, get_table_content_as_dict, + update_table, ) from badger.gui.windows.load_data_from_run_dialog import ( BadgerLoadDataFromRunDialog, ) from badger.routine import Routine -from xopt.vocs import VOCS LABEL_WIDTH = 96 @@ -172,7 +173,7 @@ def load_from_dialog(self) -> None: if not vocs.variable_names or not vocs.objective_names: dialog = QMessageBox( - text=str("Select Environment + VOCS before adding data!"), + text="Select Environment + VOCS before adding data!", parent=self, ) dialog.setIcon(QMessageBox.Information) @@ -300,7 +301,7 @@ def load_data_from_dialog(self, routine: Routine) -> None: # This happens here if selected VOCS have been changed but old data is still in the table. if self.has_data and set(data_keys) != set(filtered_table_keys): dialog = QMessageBox( - text=str( + text=( "Keys in loaded data do not match current table!" "\nTry clearing the table before adding new data." ), diff --git a/src/badger/gui/components/data_table.py b/src/badger/gui/components/data_table.py index 59b927f5..d962441c 100644 --- a/src/badger/gui/components/data_table.py +++ b/src/badger/gui/components/data_table.py @@ -3,9 +3,8 @@ copy and alternating-row styling.""" from pandas import DataFrame -from PyQt5.QtWidgets import QApplication, QTableWidget, QTableWidgetItem from PyQt5.QtCore import Qt - +from PyQt5.QtWidgets import QApplication, QTableWidget, QTableWidgetItem stylesheet = """ QTableWidget diff --git a/src/badger/gui/components/editable_table.py b/src/badger/gui/components/editable_table.py index 66d397c3..84397cf2 100644 --- a/src/badger/gui/components/editable_table.py +++ b/src/badger/gui/components/editable_table.py @@ -21,8 +21,9 @@ """ import logging +from collections.abc import Callable from functools import partial, wraps -from typing import Any, Callable, Dict, List, ParamSpec, cast +from typing import Any, ParamSpec, cast from pyparsing import TypeVar from PyQt5.QtCore import QRegExp, Qt, pyqtSignal @@ -86,9 +87,9 @@ def __init__(self, *args: Any, **kwargs: Any) -> None: header.setSectionResizeMode(1, QHeaderView.ResizeMode.Stretch) self.setColumnWidth(0, 20) # width for checkboxes - self.data: List[Dict[str, Any]] = [] - self.status: Dict[str, bool] = {} # track selection - self.formulas: Dict[str, Dict[str, Any]] = {} # track formula item + self.data: list[dict[str, Any]] = [] + self.status: dict[str, bool] = {} # track selection + self.formulas: dict[str, dict[str, Any]] = {} # track formula item self.show_selected_only = False self.keyword = "" @@ -105,7 +106,7 @@ def config_logic(self) -> None: self.itemChanged.connect(self.on_edit_table_item) def update_vocs(self) -> None: - logging.debug("Emitting data_changed signal from editable_table") + logger.debug("Emitting data_changed signal from editable_table") self.data_changed.emit() def default_info(self) -> list[Any]: @@ -262,7 +263,7 @@ def dropEvent(self, event: QDropEvent | None) -> None: event.ignore() @property - def item_names(self) -> List[str]: + def item_names(self) -> list[str]: """ Get the names of all items in the table. @@ -430,7 +431,7 @@ def add_formula_item(self, formula_tuple: tuple[str, str, dict[str, str]]) -> No def add_plain_item(self, name: str): self.add_formula_item((name, "", {})) - def get_visible_items(self) -> List[str]: + def get_visible_items(self) -> list[str]: """ Get a list of visible item names based on the current keyword filter. @@ -550,7 +551,7 @@ def add_empty_row(self): item.setForeground(QColor("gray")) self.setItem(row, 1, item) - def get_item_by_name(self, name: str) -> Dict[str, Any]: + def get_item_by_name(self, name: str) -> dict[str, Any]: """ Retrieve the item dictionary from the items list that matches the given name. @@ -688,7 +689,7 @@ def update_items_wrapper( # Add an empty row for new constraints self.add_empty_row() - def export_data(self) -> List[Dict[str, Any]]: + def export_data(self) -> list[dict[str, Any]]: """ Export the items as a list of dictionaries. diff --git a/src/badger/gui/components/eliding_label.py b/src/badger/gui/components/eliding_label.py index a7a8c865..1bd890b7 100644 --- a/src/badger/gui/components/eliding_label.py +++ b/src/badger/gui/components/eliding_label.py @@ -1,7 +1,7 @@ """QLabel that truncates text with "..." when it doesn't fit. Used for long routine names and status messages in tight layouts.""" -from PyQt5 import QtCore, QtWidgets, QtGui +from PyQt5 import QtCore, QtGui, QtWidgets class ElidingLabel(QtWidgets.QLabel): @@ -106,16 +106,16 @@ def paintEvent(self, event): class SimpleElidedLabel(QtWidgets.QLabel): def __init__(self, text="", parent=None): - super(SimpleElidedLabel, self).__init__(parent) + super().__init__(parent) self._text = text def setText(self, text): self._text = text - super(SimpleElidedLabel, self).setText(self.elidedText()) + super().setText(self.elidedText()) def resizeEvent(self, event): - super(SimpleElidedLabel, self).setText(self.elidedText()) - super(SimpleElidedLabel, self).resizeEvent(event) + super().setText(self.elidedText()) + super().resizeEvent(event) def elidedText(self): metrics = QtGui.QFontMetrics(self.font()) diff --git a/src/badger/gui/components/env_cbox.py b/src/badger/gui/components/env_cbox.py index 0fc1ed5b..8342d717 100644 --- a/src/badger/gui/components/env_cbox.py +++ b/src/badger/gui/components/env_cbox.py @@ -153,11 +153,15 @@ def __init__( self, env_dict: dict[str, Any], parent: QWidget | None = None, - envs: list[str] = [], + envs: list[str] | None = None, ): + if envs is None: + envs = [] + super().__init__(parent) self.envs = envs + self.env_dict = env_dict self.init_ui() diff --git a/src/badger/gui/components/extension_utilities.py b/src/badger/gui/components/extension_utilities.py index 79314cbe..b3c28869 100644 --- a/src/badger/gui/components/extension_utilities.py +++ b/src/badger/gui/components/extension_utilities.py @@ -4,9 +4,10 @@ import logging import time import traceback +from collections.abc import Callable from functools import wraps from types import TracebackType -from typing import Any, Callable, Optional, ParamSpec +from typing import Any, ParamSpec import matplotlib.pyplot as plt import pandas as pd @@ -55,7 +56,7 @@ def clear_tabs(tab_widget: QTabWidget) -> None: def requires_update( - last_updated: Optional[float], interval: int = 1000, requires_rebuild: bool = False + last_updated: float | None, interval: int = 1000, requires_rebuild: bool = False ) -> bool: # Check if the plot was updated recently if last_updated is not None and not requires_rebuild: @@ -73,11 +74,11 @@ def requires_update( def to_precision_float(value: Any, precision: int = 4) -> float: try: return float(f"{value:.{precision}g}") - except Exception: + except (ValueError, TypeError) as e: raise HandledException( ValueError, f"Value {value} cannot be converted to float with precision {precision}", - ) + ) from e def get_latest_reference_points( @@ -144,8 +145,8 @@ def __enter__(self) -> tuple[Figure, Axes]: def __exit__( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - exc_traceback: Optional[TracebackType], + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + exc_traceback: TracebackType | None, ) -> None: plt.close(self.fig) diff --git a/src/badger/gui/components/extensions_palette.py b/src/badger/gui/components/extensions_palette.py index 5f0accd5..055bacd9 100644 --- a/src/badger/gui/components/extensions_palette.py +++ b/src/badger/gui/components/extensions_palette.py @@ -147,7 +147,7 @@ def add_pf_viewer(self) -> None: ) except HandledException as e: QMessageBox.critical(self, "Handled Exception Error", str(e)) - except Exception: + except Exception: # noqa: BLE001 - explicit unhandled-exception fallback QMessageBox.critical( self, "Unhandled Exception Error", traceback.format_exc() ) @@ -179,7 +179,7 @@ def add_bo_visualizer(self) -> None: ) except HandledException as e: QMessageBox.critical(self, "Handled Exception Error", str(e)) - except Exception: + except Exception: # noqa: BLE001 - explicit unhandled-exception fallback QMessageBox.critical( self, "Unhandled Exception Error", traceback.format_exc() ) @@ -211,7 +211,7 @@ def add_bax_visualizer(self) -> None: ) except HandledException as e: QMessageBox.critical(self, "Handled Exception Error", str(e)) - except Exception: + except Exception: # noqa: BLE001 - explicit unhandled-exception fallback QMessageBox.critical( self, "Unhandled Exception Error", traceback.format_exc() ) @@ -237,7 +237,7 @@ def add_child_window_to_monitor(self, child_window: AnalysisExtension) -> None: self.update_palette() except HandledException as e: QMessageBox.critical(self, "Handled Exception Error", str(e)) - except Exception: + except Exception: # noqa: BLE001 - explicit unhandled-exception fallback QMessageBox.critical( self, "Unhandled Exception Error", traceback.format_exc() ) diff --git a/src/badger/gui/components/filter_cbox.py b/src/badger/gui/components/filter_cbox.py index f04413ac..6ed9857d 100644 --- a/src/badger/gui/components/filter_cbox.py +++ b/src/badger/gui/components/filter_cbox.py @@ -1,8 +1,15 @@ """Combo-box filters for narrowing the routine list by objective or region, shown in the navigator sidebar.""" -from PyQt5.QtWidgets import QVBoxLayout, QHBoxLayout, QWidget -from PyQt5.QtWidgets import QComboBox, QStyledItemDelegate, QLabel +from PyQt5.QtWidgets import ( + QComboBox, + QHBoxLayout, + QLabel, + QStyledItemDelegate, + QVBoxLayout, + QWidget, +) + from badger.gui.components.collapsible_box import CollapsibleBox diff --git a/src/badger/gui/components/generator_cbox.py b/src/badger/gui/components/generator_cbox.py index 8376f78a..a4bf3629 100644 --- a/src/badger/gui/components/generator_cbox.py +++ b/src/badger/gui/components/generator_cbox.py @@ -2,28 +2,41 @@ and configure its parameters via the Pydantic tree editor.""" from PyQt5.QtWidgets import ( - QVBoxLayout, + QCheckBox, + QComboBox, QHBoxLayout, + QLabel, + QPlainTextEdit, QPushButton, + QStyledItemDelegate, + QVBoxLayout, QWidget, - QPlainTextEdit, ) -from PyQt5.QtWidgets import QComboBox, QCheckBox, QStyledItemDelegate, QLabel +from xopt.vocs import VOCS + from badger.gui.components.collapsible_box import CollapsibleBox from badger.gui.components.pydantic_editor import BadgerPydanticEditor -from badger.settings import init_settings from badger.gui.utils import ( MouseWheelWidgetAdjustmentGuard, NoHoverFocusComboBox, ) +from badger.settings import init_settings from badger.utils import strtobool -from xopt.vocs import VOCS LABEL_WIDTH = 96 class BadgerAlgoBox(QWidget): - def __init__(self, parent=None, generators=[], scaling_functions=[]): + def __init__( + self, + parent=None, + generators: list | None = None, + scaling_functions: list | None = None, + ): + if generators is None: + generators = [] + if scaling_functions is None: + scaling_functions = [] super().__init__(parent) self.generators = generators diff --git a/src/badger/gui/components/navigators.py b/src/badger/gui/components/navigators.py index 4b4516ea..e661684b 100644 --- a/src/badger/gui/components/navigators.py +++ b/src/badger/gui/components/navigators.py @@ -4,23 +4,25 @@ context menus (open, copy path, etc.).""" import os + +from PyQt5.QtCore import QDir, Qt, QTimer, QUrl +from PyQt5.QtGui import QCursor, QDesktopServices, QFont from PyQt5.QtWidgets import ( - QWidget, - QVBoxLayout, - QTreeView, - QTreeWidget, - QTreeWidgetItem, - QMenu, QAction, QApplication, - QToolTip, QFileSystemModel, + QMenu, + QToolTip, + QTreeView, + QTreeWidget, + QTreeWidgetItem, + QVBoxLayout, + QWidget, ) -from PyQt5.QtGui import QFont, QDesktopServices, QCursor -from PyQt5.QtCore import Qt, QUrl, QTimer, QDir + from badger.archive import get_base_run_filename, get_runs -from badger.utils import run_names_to_dict from badger.settings import init_settings +from badger.utils import run_names_to_dict class FileContextMenuBase: diff --git a/src/badger/gui/components/obj_table.py b/src/badger/gui/components/obj_table.py index b5e3ab40..107212ab 100644 --- a/src/badger/gui/components/obj_table.py +++ b/src/badger/gui/components/obj_table.py @@ -2,10 +2,11 @@ MAXIMIZE rule. Supports drag-and-drop reordering and text drops.""" from typing import Any + from PyQt5.QtWidgets import ( QComboBox, - QStyledItemDelegate, QMessageBox, + QStyledItemDelegate, QWidget, ) diff --git a/src/badger/gui/components/obs_table.py b/src/badger/gui/components/obs_table.py index 69dcd8a8..8c8730c5 100644 --- a/src/badger/gui/components/obs_table.py +++ b/src/badger/gui/components/obs_table.py @@ -2,6 +2,7 @@ optimized. Supports drag-and-drop reordering and text drops.""" from typing import Any + from PyQt5.QtWidgets import ( QMessageBox, ) diff --git a/src/badger/gui/components/pf_viewer/pf_widget.py b/src/badger/gui/components/pf_viewer/pf_widget.py index ee5ef638..0af876bd 100644 --- a/src/badger/gui/components/pf_viewer/pf_widget.py +++ b/src/badger/gui/components/pf_viewer/pf_widget.py @@ -4,7 +4,6 @@ import logging import time -from typing import Optional import matplotlib.pyplot as plt import pandas as pd @@ -71,9 +70,9 @@ class ParetoFrontWidget(AnalysisWidget): parameters: ConfigurableOptions = DEFAULT_PARAMETERS # type: ignore hypervolume_history: pd.DataFrame = pd.DataFrame() - pf_1: Optional[Tensor] = None - pf_2: Optional[Tensor] = None - pf_mask: Optional[Tensor] = None + pf_1: Tensor | None = None + pf_2: Tensor | None = None + pf_mask: Tensor | None = None plot_size: tuple[float, float] = (8, 6) # UI component references @@ -87,7 +86,7 @@ class ParetoFrontWidget(AnalysisWidget): def __init__( self, routine: Routine, - parent: Optional[QWidget] = None, + parent: QWidget | None = None, ): super().__init__(routine, parent) diff --git a/src/badger/gui/components/plot_event_handlers.py b/src/badger/gui/components/plot_event_handlers.py index 6a80a822..73ac038b 100644 --- a/src/badger/gui/components/plot_event_handlers.py +++ b/src/badger/gui/components/plot_event_handlers.py @@ -1,24 +1,22 @@ """Mouse event handlers for matplotlib plots in the analysis extensions. Handles click, hover, scroll-zoom, and data-point annotation tooltips.""" -from badger.routine import Routine -from matplotlib.backend_bases import MouseEvent, MouseButton, PickEvent +import logging +from typing import cast + +import pandas as pd +from matplotlib.backend_bases import MouseButton, MouseEvent, PickEvent from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg +from matplotlib.collections import PathCollection +from matplotlib.text import Annotation +from pyparsing import Callable -from badger.gui.components.types import ConfigurableOptions from badger.gui.components.extension_utilities import ( HandledException, to_precision_float, ) - -from typing import cast - -import logging - -from matplotlib.collections import PathCollection -from matplotlib.text import Annotation -import pandas as pd -from pyparsing import Callable +from badger.gui.components.types import ConfigurableOptions +from badger.routine import Routine logger = logging.getLogger(__name__) @@ -356,8 +354,8 @@ def on_pick(self, event: PickEvent) -> None: xy=point, xytext=(0, 0), # Initial position of the tooltip textcoords="offset pixels", - bbox=dict(boxstyle="round", fc="w"), - arrowprops=dict(arrowstyle="->"), + bbox={"boxstyle": "round", "fc": "w"}, + arrowprops={"arrowstyle": "->"}, ) # Adjust tooltip position based on the region and the size of the text diff --git a/src/badger/gui/components/process_manager.py b/src/badger/gui/components/process_manager.py index ddfebeb9..8beee914 100644 --- a/src/badger/gui/components/process_manager.py +++ b/src/badger/gui/components/process_manager.py @@ -1,9 +1,7 @@ """Keeps a pool of pre-spawned subprocesses ready to run optimizations. When one is consumed by a run, signals that a new one should be created.""" -from typing import Dict, Optional - -from PyQt5.QtCore import pyqtSignal, QObject +from PyQt5.QtCore import QObject, pyqtSignal class ProcessManager(QObject): @@ -18,7 +16,7 @@ def __init__(self) -> None: super().__init__() self.processes_queue = [] - def add_to_queue(self, process_with_args: Dict) -> None: + def add_to_queue(self, process_with_args: dict) -> None: """ Add to a dict contaitng a process and it's coresponding args to the processes_queue. @@ -28,7 +26,7 @@ def add_to_queue(self, process_with_args: Dict) -> None: """ self.processes_queue.append(process_with_args) - def remove_from_queue(self) -> Optional[Dict]: + def remove_from_queue(self) -> dict | None: """ Removes and returns a process and it's coresponding args to the processes_queue. If no process are in the processes_queue then the method returns None. @@ -52,7 +50,7 @@ def close_proccesses(self) -> bool: ------- True: bool """ - for i in range(0, len(self.processes_queue)): + for i in range(len(self.processes_queue)): process = self.processes_queue.pop(0) process["process"].terminate() process["process"].join() diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index 1a7f49e1..dd29ba52 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -10,15 +10,15 @@ import ast import logging import re +from collections.abc import Callable, Sequence from dataclasses import dataclass from inspect import isclass from types import NoneType from typing import ( Annotated, Any, - Callable, + ClassVar, Optional, - Sequence, TypeVar, Union, cast, @@ -27,6 +27,8 @@ ) import yaml +from bax_algorithms.emittance import PathwiseMinimizeEmittance +from bax_algorithms.solenoid_alignment import PathwiseSolenoidAlignment from pydantic import BaseModel, Field, ValidationError, create_model from pydantic.fields import FieldInfo from pydantic_core import PydanticUndefined, PydanticUndefinedType @@ -47,19 +49,16 @@ QVBoxLayout, QWidget, ) +from torch import Tensor from xopt.errors import VOCSError from xopt.generators import get_generator from xopt.generators.bayesian.bax.algorithms import Algorithm from xopt.generators.bayesian.bax_generator import BaxGenerator from xopt.generators.bayesian.bayesian_generator import BayesianGenerator from xopt.generators.bayesian.turbo import TurboController -from torch import Tensor from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS -from bax_algorithms.emittance import PathwiseMinimizeEmittance -from bax_algorithms.solenoid_alignment import PathwiseSolenoidAlignment - logger = logging.getLogger(__name__) @@ -80,8 +79,8 @@ def tuple_constructor(loader: Any, node: yaml.ScalarNode) -> Any: if TUPLE_PATTERN.match(value): try: return ast.literal_eval(value) - except Exception: - pass + except (ValueError, SyntaxError): + logger.warning(f"Failed to parse tuple from string: {value}") return value @@ -96,10 +95,10 @@ def convert_to_type(value: Any, type: Callable[[Any], T]) -> T: def _set_value_for_basic_widget( widget: QWidget, - value: str | float | int | bool | None, + value: str | float | bool | None, ) -> None: nullable = bool(widget.property("badger_nullable")) - if isinstance(widget, QLabel) or isinstance(widget, QLineEdit): + if isinstance(widget, (QLabel, QLineEdit)): widget.setText("null" if value is None else str(value)) elif isinstance(widget, QDoubleSpinBox): if value is None and nullable: @@ -151,10 +150,8 @@ def find_primary( return resolved[0] @classmethod - def resolve( - cls, annotation: type[Any] | Union[Any, None] | None - ) -> "BadgerResolvedType": - origin: type[Any] | Union[Any, None] | None = get_origin(annotation) + def resolve(cls, annotation: type[Any] | Any | None) -> "BadgerResolvedType": + origin: type[Any] | Any | None = get_origin(annotation) args = get_args(annotation) nullable = False @@ -217,8 +214,8 @@ def resolve( @classmethod def resolve_qt( cls, - annotation: type[Any] | Union[Any, None] | None, - default: float | int | bool | dict[str, Any] | list[Any] | None = None, + annotation: type[Any] | Any | None, + default: float | bool | dict[str, Any] | list[Any] | None = None, editor_info: tuple["BadgerPydanticEditor", QTreeWidgetItem] | None = None, ) -> QWidget | None: resolved_type = BadgerResolvedType.resolve(annotation) @@ -229,11 +226,11 @@ def resolve_qt( widget = QLineEdit() widget.setText("null") elif issubclass(resolved_type.main, BaseModel): - if issubclass(resolved_type.main, TurboController): - widget = QComboBox() - elif issubclass(resolved_type.main, NumericalOptimizer): - widget = QComboBox() - elif issubclass(resolved_type.main, Algorithm): + if ( + issubclass(resolved_type.main, TurboController) + or issubclass(resolved_type.main, NumericalOptimizer) + or issubclass(resolved_type.main, Algorithm) + ): widget = QComboBox() else: return None @@ -242,9 +239,12 @@ def resolve_qt( if default is None: default = {"name": "null"} - if isinstance(default, dict) and "name" in default: - if (index := widget.findText(default["name"])) >= 0: - widget.setCurrentIndex(index) + if ( + isinstance(default, dict) + and "name" in default + and (index := widget.findText(default["name"])) >= 0 + ): + widget.setCurrentIndex(index) elif resolved_type.main == NoneType: widget = QLabel() widget.setText("null") @@ -395,7 +395,7 @@ def _qt_widget_to_yaml_value(widget: Any) -> str | None: return None elif isinstance(widget, BadgerListEditor): return widget.get_parameters_yaml() - elif isinstance(widget, QSpinBox) or isinstance(widget, QDoubleSpinBox): + elif isinstance(widget, (QSpinBox, QDoubleSpinBox)): if widget.property("badger_nullable") and widget.value() == widget.minimum(): return "null" return str(widget.value()) @@ -458,7 +458,7 @@ def _qt_widget_to_value(widget: Any) -> Any: return None elif isinstance(widget, BadgerListEditor): return widget.get_parameters_dict() - elif isinstance(widget, QSpinBox) or isinstance(widget, QDoubleSpinBox): + elif isinstance(widget, (QSpinBox, QDoubleSpinBox)): if widget.property("badger_nullable") and widget.value() == widget.minimum(): return None return widget.value() @@ -665,7 +665,7 @@ def get_parameters_dict(self) -> dict[str, Any] | None: class BadgerPydanticEditor(QTreeWidget): vocs: VOCS = VOCS(variables={}) - defaults: dict[str, Any] = {} + generator_name: str = "" model_class: type[BaseModel] | None = None @@ -680,7 +680,7 @@ class BadgerPydanticEditor(QTreeWidget): # the generator's ``name`` field). The effective set is the union of both, # resolved by ``get_excluded_fields``. COMMON_EXCLUDED_FIELDS: frozenset[str] = frozenset({"computation_time"}) - GENERATOR_EXCLUDED_FIELDS: dict[str, frozenset[str]] = { + GENERATOR_EXCLUDED_FIELDS: ClassVar[dict[str, frozenset[str]]] = { # "bax": frozenset({"algorithm_results"}), } @@ -710,9 +710,10 @@ def __init__( value_col: int = 1, update_callback: Callable[["BadgerPydanticEditor"], None] | None = None, ): + self.defaults: dict[str, Any] = {} + QTreeWidget.__init__(self, parent) - if value_col < 1: - value_col = 1 + value_col = max(value_col, 1) self.value_col = value_col self.update_callback = update_callback self.setColumnCount(self.value_col + 1) @@ -720,7 +721,7 @@ def __init__( self.setHeaderLabels( [ "Parameter" if i == 0 else "Value" if i == self.value_col else "" - for i in range(0, self.value_col + 1) + for i in range(self.value_col + 1) ] ) @@ -728,14 +729,14 @@ def __init__( def _set_params_recurse( self, - parent: Optional[QTreeWidgetItem], + parent: QTreeWidgetItem | None, fields: dict[str, FieldInfo], defaults: dict[str, Any] | None, hidden: bool, ) -> None: for field_name, field_info in fields.items(): child = QTreeWidgetItem( - [field_name if i == 0 else "" for i in range(0, self.value_col + 1)] + [field_name if i == 0 else "" for i in range(self.value_col + 1)] ) if parent is None: @@ -1078,13 +1079,18 @@ def update_params_from_generator_class( @staticmethod def filter_class_fields( pydantic_class: type[BaseModel], - fields_to_remove: list[str] = [], - defaults: dict[str, Any] = {}, + fields_to_remove: list[str] | None = None, + defaults: dict[str, Any] | None = None, include_defaults: bool = False, excluded_fields: frozenset[str] = frozenset(), ) -> tuple[dict[str, FieldInfo], dict[str, FieldInfo]]: condition: Callable[[str], bool] + if fields_to_remove is None: + fields_to_remove = [] + if defaults is None: + defaults = {} + def include_condition(k: str) -> bool: return k in defaults and k not in fields_to_remove @@ -1113,7 +1119,7 @@ def exclude_condition(k: str) -> bool: @staticmethod def get_defaults_from_type(pydantic_class: type[Any]) -> dict[str, Any]: if not issubclass(pydantic_class, BaseModel): - raise ValueError("Provided class is not a Pydantic model") + raise TypeError("Provided class is not a Pydantic model") defaults: dict[str, Any] = {} for field_name, field_info in pydantic_class.model_fields.items(): if field_info.default is not PydanticUndefined: @@ -1263,11 +1269,11 @@ def _inject_computed_fields( continue try: parameters_dict[cf_name] = getattr(instance, cf_name) - except Exception as e: + except Exception as e: # noqa: BLE001 - computed field runs model code logger.debug( f"Could not compute {cf_name} on {model_class.__name__}: {e}" ) - except Exception as e: + except Exception as e: # noqa: BLE001 - model_construct runs model code logger.debug( f"Could not model_construct {model_class.__name__} for computed-field injection: {e}" ) @@ -1299,7 +1305,7 @@ def convert_dict(val: Any) -> Any: # Convert str-encoded dicts (and lists) back into actual dict objects.""" if isinstance(val, str): stripped = val.strip() - if stripped.startswith("{") or stripped.startswith("["): + if stripped.startswith(("{", "[")): try: return ast.literal_eval(stripped) except (ValueError, SyntaxError): @@ -1350,9 +1356,7 @@ def convert_dict(val: Any) -> Any: self.update_error_styles(loc, msg) def update_error_styles(self, loc: tuple[int | str, ...], msg: str) -> None: - error_widget: QTreeWidgetItem | QTreeWidget | "BadgerPydanticEditor" | None = ( - None - ) + error_widget: QTreeWidgetItem | QTreeWidget | BadgerPydanticEditor | None = None if len(loc) > 0: error_widget = self.find_widget_at_path(loc) else: diff --git a/src/badger/gui/components/reorderable_table.py b/src/badger/gui/components/reorderable_table.py index 99804c28..a76be5fd 100644 --- a/src/badger/gui/components/reorderable_table.py +++ b/src/badger/gui/components/reorderable_table.py @@ -8,7 +8,7 @@ # https://creativecommons.org/publicdomain/zero/1.0/ # https://creativecommons.org/publicdomain/zero/1.0/legalcode -from PyQt5 import QtWidgets, QtGui +from PyQt5 import QtGui, QtWidgets class MyModel(QtGui.QStandardItemModel): @@ -55,7 +55,7 @@ def __init__(self, parent): self.setModel(self.model) for idx, data in enumerate(["foo", "bar", "baz"]): - item_1 = QtGui.QStandardItem("Item {}".format(idx)) + item_1 = QtGui.QStandardItem(f"Item {idx}") item_1.setEditable(False) item_1.setDropEnabled(False) diff --git a/src/badger/gui/components/robust_spinbox.py b/src/badger/gui/components/robust_spinbox.py index 918910df..706c1988 100644 --- a/src/badger/gui/components/robust_spinbox.py +++ b/src/badger/gui/components/robust_spinbox.py @@ -2,8 +2,9 @@ and optional read-only mode. Used for variable bounds and constraint thresholds.""" -from PyQt5.QtWidgets import QDoubleSpinBox, QAbstractSpinBox from PyQt5.QtCore import Qt +from PyQt5.QtWidgets import QAbstractSpinBox, QDoubleSpinBox + from badger.gui.utils import MouseWheelWidgetAdjustmentGuard @@ -12,7 +13,7 @@ def __init__(self, *args, **kwargs): try: decimals = kwargs["decimals"] del kwargs["decimals"] - except: + except KeyError: decimals = 6 try: @@ -20,7 +21,7 @@ def __init__(self, *args, **kwargs): del kwargs["lower_bound"] if lb is None: lb = -1e3 - except: + except KeyError: lb = -1e3 try: @@ -28,7 +29,7 @@ def __init__(self, *args, **kwargs): del kwargs["upper_bound"] if ub is None: ub = 1e3 - except: + except KeyError: ub = 1e3 try: @@ -36,7 +37,7 @@ def __init__(self, *args, **kwargs): del kwargs["default_value"] if default_value is None: default_value = 0 - except: + except KeyError: default_value = 0 super().__init__(*args, **kwargs) diff --git a/src/badger/gui/components/routine_editor.py b/src/badger/gui/components/routine_editor.py index b03da465..7482e4df 100644 --- a/src/badger/gui/components/routine_editor.py +++ b/src/badger/gui/components/routine_editor.py @@ -1,12 +1,19 @@ """Container that wraps the routine page in a scrollable area with save/cancel/delete buttons. Coordinates creation, editing, and deletion.""" -from PyQt5.QtWidgets import QVBoxLayout, QHBoxLayout, QWidget, QPushButton -from PyQt5.QtWidgets import QTextEdit, QStackedWidget, QScrollArea from PyQt5.QtCore import pyqtSignal from PyQt5.QtGui import QFont -from badger.gui.components.routine_page import BadgerRoutinePage +from PyQt5.QtWidgets import ( + QHBoxLayout, + QPushButton, + QScrollArea, + QStackedWidget, + QTextEdit, + QVBoxLayout, + QWidget, +) +from badger.gui.components.routine_page import BadgerRoutinePage from badger.routine import Routine diff --git a/src/badger/gui/components/routine_item.py b/src/badger/gui/components/routine_item.py index e24f22db..2343b267 100644 --- a/src/badger/gui/components/routine_item.py +++ b/src/badger/gui/components/routine_item.py @@ -2,10 +2,18 @@ and environment with hover/selection styling and a delete button.""" from datetime import datetime -from PyQt5.QtWidgets import QWidget, QHBoxLayout, QVBoxLayout, QLabel -from PyQt5.QtWidgets import QSizePolicy, QMessageBox + from PyQt5.QtCore import Qt, pyqtSignal from PyQt5.QtGui import QFont +from PyQt5.QtWidgets import ( + QHBoxLayout, + QLabel, + QMessageBox, + QSizePolicy, + QVBoxLayout, + QWidget, +) + from badger.gui.components.eliding_label import ElidingLabel from badger.gui.utils import create_button diff --git a/src/badger/gui/components/routine_page.py b/src/badger/gui/components/routine_page.py index 7f5610bf..024f272e 100644 --- a/src/badger/gui/components/routine_page.py +++ b/src/badger/gui/components/routine_page.py @@ -30,7 +30,7 @@ import os import traceback import warnings -from datetime import datetime +from datetime import UTC, datetime from functools import partial from typing import Any @@ -41,11 +41,11 @@ from gest_api.vocs import ( BaseConstraint, BaseObjective, + ContinuousVariable, GreaterThanConstraint, LessThanConstraint, MaximizeObjective, MinimizeObjective, - ContinuousVariable, ) from pydantic import ValidationError from PyQt5.QtCore import Qt, QTimer, pyqtSignal @@ -92,7 +92,7 @@ from badger.gui.components.env_cbox import BadgerEnvBox from badger.gui.components.filter_cbox import BadgerFilterBox from badger.gui.components.generator_cbox import BadgerAlgoBox -from badger.gui.utils import filter_generator_config +from badger.gui.utils import filter_generator_config, with_busy_cursor from badger.gui.windows.add_random_dialog import BadgerAddRandomDialog from badger.gui.windows.docs_window import BadgerDocsWindow from badger.gui.windows.edit_script_dialog import BadgerEditScriptDialog @@ -101,7 +101,6 @@ ) from badger.gui.windows.lim_vrange_dialog import BadgerLimitVariableRangeDialog from badger.gui.windows.message_dialog import BadgerScrollableMessageBox -from badger.gui.utils import with_busy_cursor from badger.gui.windows.review_dialog import BadgerReviewDialog from badger.routine import Routine from badger.settings import init_settings @@ -152,7 +151,7 @@ def extract_objective_symbol(objective: BaseObjective) -> str: if isinstance(objective, MaximizeObjective): return "MAXIMIZE" else: - raise ValueError(f"Unknown objective type: {objective}") + raise TypeError(f"Unknown objective type: {objective}") class BadgerRoutinePage(QWidget): @@ -689,28 +688,27 @@ def _filter_generator_params(self, generator_name: str, generator_config: dict): Filter which generator parameters get saved to template """ - if generator_name in ["expected_improvement", "upper_confidence_bound"]: - if ( - "turbo_controller" in generator_config - and generator_config["turbo_controller"] is not None - and isinstance(generator_config["turbo_controller"], dict) - ): - turbo = generator_config["turbo_controller"] - generator_config["turbo_controller"] = { - k: v - for k, v in turbo.items() - if k - in { - "name", - "length", - "length_max", - "length_min", - "failure_tolerance", - "success_tolerance", - "scale_factor", - "restrict_model_data", - } + if generator_name in ["expected_improvement", "upper_confidence_bound"] and ( + "turbo_controller" in generator_config + and generator_config["turbo_controller"] is not None + and isinstance(generator_config["turbo_controller"], dict) + ): + turbo = generator_config["turbo_controller"] + generator_config["turbo_controller"] = { + k: v + for k, v in turbo.items() + if k + in { + "name", + "length", + "length_max", + "length_min", + "failure_tolerance", + "success_tolerance", + "scale_factor", + "restrict_model_data", } + } return generator_config @@ -726,7 +724,7 @@ def save_template_yaml(self): # Suggest a filename based on the routine name or placeholder routine_name = self.edit_save.text() or self.edit_save.placeholderText() if not routine_name: - routine_name = "template_" + datetime.now().strftime("%y%m%d_%H%M%S") + routine_name = "template_" + datetime.now(tz=UTC).strftime("%y%m%d_%H%M%S") suggested_filename = f"{routine_name}.yaml" template_path, _ = QFileDialog.getSaveFileName( self, @@ -1006,7 +1004,7 @@ def refresh_ui(self, routine: Routine | None = None, silent: bool = False) -> No self.edit_save.setText(routine.name) self.edit_descr.setPlainText(routine.description) - self.generator_box.check_use_script.setChecked(not not self.script) + self.generator_box.check_use_script.setChecked(bool(self.script)) def set_routine(self, routine: Routine, silent: bool = False): self.refresh_ui(routine, silent=silent) @@ -1039,7 +1037,7 @@ def select_generator(self, i: int): # Get vocs try: vocs, _ = self.env_box.compose_vocs() - except Exception: + except BadgerRoutineError: vocs = None self.generator_box.edit.set_params_from_generator(name, filtered_config, vocs) @@ -1090,7 +1088,9 @@ def create_env(self): try: env = instantiate_env(self.env, configs) except Exception as e: - raise BadgerEnvInstantiationError(f"Failed to instantiate environment: {e}") + raise BadgerEnvInstantiationError( + f"Failed to instantiate environment: {e}" + ) from e return env @@ -1100,10 +1100,11 @@ def refresh_params_generator(self): try: tmp = {} - exec(self.script, tmp) + # User-provided script must define a `generate` function, so exec is required here. + exec(self.script, tmp) # noqa: S102 - runs user-provided generator script try: tmp["generate"] # test if generate function is defined - except Exception as e: + except KeyError as e: QMessageBox.warning( self, "Please define a valid generate function!", str(e) ) @@ -1113,14 +1114,14 @@ def refresh_params_generator(self): # Get vocs try: vocs, _ = self.env_box.compose_vocs() - except Exception: + except BadgerRoutineError: vocs = None # Function generate comes from the script params_generator = tmp["generate"](env, vocs) self.generator_box.edit.set_params_from_generator( self.routine.generator.name, params_generator, vocs ) - except Exception as e: + except Exception as e: # noqa: BLE001 - explicit unhandled-exception fallback QMessageBox.warning(self, "Invalid script!", str(e)) @with_busy_cursor @@ -1167,7 +1168,7 @@ def select_env(self, i: int): self.env_box.btn_refresh.setDisabled(False) if self.generator_box.check_use_script.isChecked(): self.refresh_params_generator() - except Exception: + except Exception: # noqa: BLE001 - env selection rollback on any failure self.configs = None self.env = None self.env_box.cb.setCurrentIndex(-1) @@ -1285,7 +1286,7 @@ def fill_curr_in_init_table(self, record=False): raise BadgerEnvVarError( f"Failed to get current variable values : {e}\n" "Please ensure the environment is properly configured." - ) + ) from e # Iterate through the rows for row in range(table.rowCount()): @@ -1320,7 +1321,7 @@ def add_rand_in_init_table(self, add_rand_config=None, record=True): # get small region around current point to sample try: vocs, _ = self.env_box.compose_vocs() - except Exception: + except BadgerRoutineError: # Switch to manual mode to allow the user fixing the vocs issue QMessageBox.warning( self, @@ -1490,7 +1491,7 @@ def set_vrange(self, set_all=True): raise BadgerEnvVarError( f"Failed to get current variable values : {e}\n" "Please ensure the environment is properly configured." - ) + ) from e option_idx = self.limit_option["limit_option_idx"] # 0: ratio with current value, 1: ratio with full range, 2: delta around current value @@ -1527,7 +1528,7 @@ def set_vrange(self, set_all=True): self.update_init_table() # auto populate if option is set # remember user selection for applying limit changes - if not self.lim_apply_to_vars == 2: + if self.lim_apply_to_vars != 2: # Check if lim_apply_to_vars has been initialized # It will be set to 2 until the btn_lim_vrange is clicked self.lim_apply_to_vars = set_all @@ -1732,7 +1733,7 @@ def toggle_relative_to_curr(self, checked, refresh=True): if checked: try: _ = self.env_box.compose_vocs() - except Exception: + except BadgerRoutineError: logger.warning("Variable range is not valid, switching to manual mode.") QTimer.singleShot(0, lambda: self.env_box.relative_to_curr.click()) QMessageBox.warning( @@ -1895,10 +1896,9 @@ def _compose_routine(self) -> Routine: NO_OBJECTIVE_GENERATORS = ["bax"] - if not vocs.objectives: - if generator_name not in NO_OBJECTIVE_GENERATORS: - logger.error("No objectives selected.") - raise BadgerRoutineError("no objectives selected") + if not vocs.objectives and generator_name not in NO_OBJECTIVE_GENERATORS: + logger.error("No objectives selected.") + raise BadgerRoutineError("no objectives selected") # Initial points init_points_df = pd.DataFrame.from_dict( @@ -1957,7 +1957,9 @@ def _compose_routine(self) -> Routine: # Metadata badger_version=get_badger_version(), xopt_version=get_xopt_version(), - creation_ts=ts_float_to_str(datetime.now().timestamp(), "lcls-fname"), + creation_ts=ts_float_to_str( + datetime.now(tz=UTC).timestamp(), "lcls-fname" + ), # Xopt part generator=generator, # Badger part @@ -1996,7 +1998,7 @@ def _compose_routine(self) -> Routine: def review(self): try: routine = self._compose_routine() - except Exception: + except Exception: # noqa: BLE001 - explicit unhandled-exception fallback return QMessageBox.critical( self, "Invalid routine!", traceback.format_exc() ) @@ -2016,7 +2018,7 @@ def update_description(self): "Update success!", f"Routine {self.routine.name} description was updated!", ) - except Exception: + except Exception: # noqa: BLE001 - catch all exceptions during routine update return QMessageBox.critical(self, "Update failed!", traceback.format_exc()) def set_default_generator(self, generator_name: str) -> None: diff --git a/src/badger/gui/components/routine_runner.py b/src/badger/gui/components/routine_runner.py index c37a4f37..2efbecb7 100644 --- a/src/badger/gui/components/routine_runner.py +++ b/src/badger/gui/components/routine_runner.py @@ -12,30 +12,29 @@ import traceback import pandas as pd -from PyQt5.QtCore import pyqtSignal, QObject, QTimer +from PyQt5.QtCore import QObject, QTimer, pyqtSignal from PyQt5.QtWidgets import QDialog from badger.errors import ( - BadgerRunTerminated, - BadgerError, - MEASUREMENT_ERROR_TYPE, - MEASUREMENT_ACTION_TYPE, - MEASUREMENT_ACTION_RETRY, MEASUREMENT_ACTION_ABORT, - TERMINATION_REACHED_TYPE, - TERMINATION_ACTION_TYPE, + MEASUREMENT_ACTION_RETRY, + MEASUREMENT_ACTION_TYPE, + MEASUREMENT_ERROR_TYPE, TERMINATION_ACTION_CONTINUE, TERMINATION_ACTION_END, + TERMINATION_ACTION_TYPE, + TERMINATION_REACHED_TYPE, + BadgerError, + BadgerRunTerminated, ) -from badger.tests.utils import get_current_vars -from badger.routine import calculate_variable_bounds, calculate_initial_points -from badger.settings import init_settings from badger.gui.components.process_manager import ProcessManager from badger.gui.windows.measurement_retry_dialog import BadgerMeasurementRetryDialog from badger.gui.windows.termination_reached_dialog import ( BadgerTerminationReachedDialog, ) -from badger.routine import Routine +from badger.routine import Routine, calculate_initial_points, calculate_variable_bounds +from badger.settings import init_settings +from badger.tests.utils import get_current_vars logger = logging.getLogger(__name__) @@ -60,7 +59,7 @@ def __init__( self, process_manager: ProcessManager, routine: Routine = None, - routine_filename: str = None, + routine_filename: str | None = None, save: bool = False, verbose: int = 2, use_full_ts: bool = False, @@ -229,7 +228,7 @@ def run(self, run_data_flag: bool = False, init_points_flag: bool = False) -> No except BadgerRunTerminated as e: self.signals.finished.emit() self.signals.info.emit(str(e)) - except Exception as e: + except Exception as e: # noqa: BLE001 - run worker boundary traceback_info = traceback.format_exc() e._details = traceback_info self.signals.finished.emit() diff --git a/src/badger/gui/components/run_monitor.py b/src/badger/gui/components/run_monitor.py index f79c846e..9a707f3c 100644 --- a/src/badger/gui/components/run_monitor.py +++ b/src/badger/gui/components/run_monitor.py @@ -11,7 +11,7 @@ import os import traceback from importlib import resources -from typing import TYPE_CHECKING, List, Optional +from typing import TYPE_CHECKING import numpy as np import pandas as pd @@ -73,7 +73,7 @@ class BadgerOptMonitor(QWidget): sig_toggle_other = pyqtSignal(bool) sig_env_ready = pyqtSignal() - def __init__(self, process_manager: "Optional[ProcessManager]" = None): + def __init__(self, process_manager: "ProcessManager | None" = None): super().__init__() # self.setAttribute(Qt.WA_DeleteOnClose, True) @@ -239,7 +239,9 @@ def config_logic(self) -> None: self.cb_plot_y.currentIndexChanged.connect(self.select_x_plot_y_axis) self.check_relative.stateChanged.connect(self.toggle_x_plot_y_axis_relative) - def init_plots(self, routine: Routine = None, run_filename: str = None) -> None: + def init_plots( + self, routine: Routine = None, run_filename: str | None = None + ) -> None: """ Initialize and configure the plots and related components in the application. @@ -289,16 +291,16 @@ def init_plots(self, routine: Routine = None, run_filename: str = None) -> None: self.monitor.removeItem(self.plot_con) self.plot_con.removeItem(self.inspector_constraint) del self.plot_con - except: - pass + except AttributeError: + logger.debug("No constraints plot to remove") # if statics exist delete that plot try: self.monitor.removeItem(self.plot_obs) self.plot_obs.removeItem(self.inspector_state) del self.plot_obs - except: - pass + except AttributeError: + logger.debug("No observables plot to remove") # if no routine is loaded set button to disabled self.sig_lock_action.emit() @@ -328,8 +330,8 @@ def init_plots(self, routine: Routine = None, run_filename: str = None) -> None: # Configure constraint plots if constraint_names: try: - self.plot_con - except: + _ = self.plot_con + except AttributeError: self.plot_con = plot_con = add_axes( self.monitor, "constraints", @@ -349,14 +351,14 @@ def init_plots(self, routine: Routine = None, run_filename: str = None) -> None: self.monitor.removeItem(self.plot_con) self.plot_con.removeItem(self.inspector_constraint) del self.plot_con - except: - pass + except AttributeError: + logger.debug("No constraints plot to remove") # Configure state plots if sta_names: try: - self.plot_obs - except: + _ = self.plot_obs + except AttributeError: self.plot_obs = plot_obs = add_axes( self.monitor, "observables", @@ -375,8 +377,8 @@ def init_plots(self, routine: Routine = None, run_filename: str = None) -> None: self.monitor.removeItem(self.plot_obs) self.plot_obs.removeItem(self.inspector_state) del self.plot_obs - except: - pass + except AttributeError: + logger.debug("No observables plot to remove") # Reset inspectors self.inspector_objective.setValue(0) @@ -667,7 +669,7 @@ def routine_finished(self) -> None: try: env.interface.stop_recording(os.path.join(path, filename)) except AttributeError: # recording was not enabled - pass + logger.debug("Recording was not enabled") self.sig_run_name.emit(run["filename"]) self.sig_status.emit( @@ -678,9 +680,9 @@ def routine_finished(self) -> None: # self, 'Success!', # f'Archive success: Run data archived to {BADGER_ARCHIVE_ROOT}') - except Exception as e: + except Exception as e: # noqa: BLE001 - archive external call self.sig_run_name.emit(None) - self.sig_status.emit(f"Archive failed: {str(e)}") + self.sig_status.emit(f"Archive failed: {e!s}") # if not self.testing: # QMessageBox.critical(self, 'Archive failed!', # f'Archive failed: {str(e)}') @@ -695,12 +697,12 @@ def destroy_unused_env(self) -> None: try: del self.routine_runner.routine.environment except AttributeError: # env already destroyed - pass + logger.debug("Environment already destroyed") try: del self.routine.environment except AttributeError: # env already destroyed - pass + logger.debug("Environment already destroyed") def on_error(self, error: Exception) -> None: details = error._details if hasattr(error, "_details") else None @@ -719,8 +721,8 @@ def on_info(self, msg) -> None: def logbook(self) -> None: try: send_to_logbook(self.routine, self.monitor) - except Exception as e: - self.sig_status.emit(f"Log failed: {str(e)}") + except Exception as e: # noqa: BLE001 - logbook external call + self.sig_status.emit(f"Log failed: {e!s}") # QMessageBox.critical(self, 'Log failed!', str(e)) return @@ -768,7 +770,7 @@ def sync_ins(self, pos) -> None: try: ts = self.extract_timestamp() value = idx = np.clip(np.round(pos), 0, len(ts) - 1) - except: # no data + except (IndexError, ValueError): # no data value = idx = np.round(pos) self.inspector_objective.setValue(value) if self.vocs and self.vocs.constraint_names: @@ -821,7 +823,7 @@ def reset_env(self) -> None: # QMessageBox.information(self, 'Reset Environment', # f'Env vars {curr_vars} -> {self.init_vars}') - def get_checkpoint(self) -> Optional[dict[str, float]]: + def get_checkpoint(self) -> dict[str, float] | None: if not self.routine or not self.routine.environment or not self.routine.vocs: return None @@ -1098,7 +1100,7 @@ def create_cursor_line() -> pg.InfiniteLine: ) -def set_data(names: List[str], curves: dict, data: pd.DataFrame, ts=None) -> None: +def set_data(names: list[str], curves: dict, data: pd.DataFrame, ts=None) -> None: # Split data into live and not live live_mask = data["live"].astype(bool) live_data = data.loc[live_mask] diff --git a/src/badger/gui/components/state_item.py b/src/badger/gui/components/state_item.py index 1edd875b..f0901a34 100644 --- a/src/badger/gui/components/state_item.py +++ b/src/badger/gui/components/state_item.py @@ -1,8 +1,8 @@ """Creates a single observable-state row: a combo box for the observable name and a remove button.""" -from PyQt5.QtWidgets import QHBoxLayout, QPushButton, QWidget -from PyQt5.QtWidgets import QStyledItemDelegate +from PyQt5.QtWidgets import QHBoxLayout, QPushButton, QStyledItemDelegate, QWidget + from badger.gui.utils import NoHoverFocusComboBox @@ -17,7 +17,7 @@ def state_item(options, remove_item, name=None): cb_sta.addItems(options) try: idx = options.index(name) - except: + except ValueError: idx = 0 cb_sta.setCurrentIndex(idx) diff --git a/src/badger/gui/components/status_bar.py b/src/badger/gui/components/status_bar.py index 15b3bfb5..2c5751a5 100644 --- a/src/badger/gui/components/status_bar.py +++ b/src/badger/gui/components/status_bar.py @@ -1,11 +1,13 @@ """Bottom status bar showing run status and a settings button.""" from importlib import resources -from PyQt5.QtWidgets import QHBoxLayout, QWidget, QPushButton + +from PyQt5.QtCore import QSize, Qt from PyQt5.QtGui import QIcon -from PyQt5.QtCore import Qt, QSize -from badger.gui.windows.settings_dialog import BadgerSettingsDialog +from PyQt5.QtWidgets import QHBoxLayout, QPushButton, QWidget + from badger.gui.components.eliding_label import SimpleElidedLabel +from badger.gui.windows.settings_dialog import BadgerSettingsDialog class BadgerStatusBar(QWidget): diff --git a/src/badger/gui/components/syntax.py b/src/badger/gui/components/syntax.py index d8195918..1c0c8885 100644 --- a/src/badger/gui/components/syntax.py +++ b/src/badger/gui/components/syntax.py @@ -1,6 +1,8 @@ """Python syntax highlighter for the script editor. Colors keywords, strings, comments, numbers, and operators.""" +from typing import ClassVar + from PyQt5 import QtCore, QtGui @@ -37,7 +39,7 @@ class PythonHighlighter(QtGui.QSyntaxHighlighter): """Syntax highlighter for the Python language.""" # Python keywords - keywords = [ + keywords: ClassVar[list[str]] = [ "and", "assert", "break", @@ -74,7 +76,7 @@ class PythonHighlighter(QtGui.QSyntaxHighlighter): ] # Python operators - operators = [ + operators: ClassVar[list[str]] = [ "=", # Comparison "==", @@ -107,7 +109,7 @@ class PythonHighlighter(QtGui.QSyntaxHighlighter): ] # Python braces - braces = [ + braces: ClassVar[list[str]] = [ "{", "}", "(", @@ -127,12 +129,10 @@ def __init__(self, parent: QtGui.QTextDocument) -> None: # Keyword, operator, and brace rules rules += [ - (r"\b%s\b" % w, 0, STYLES["keyword"]) for w in PythonHighlighter.keywords - ] - rules += [ - (r"%s" % o, 0, STYLES["operator"]) for o in PythonHighlighter.operators + (rf"\b{w}\b", 0, STYLES["keyword"]) for w in PythonHighlighter.keywords ] - rules += [(r"%s" % b, 0, STYLES["brace"]) for b in PythonHighlighter.braces] + rules += [(o, 0, STYLES["operator"]) for o in PythonHighlighter.operators] + rules += [(b, 0, STYLES["brace"]) for b in PythonHighlighter.braces] # All other rules rules += [ @@ -163,21 +163,20 @@ def highlightBlock(self, text): # Do other syntax formatting for expression, nth, format in self.rules: index = expression.indexIn(text, 0) - if index >= 0: - # if there is a string we check - # if there are some triple quotes within the string - # they will be ignored if they are matched again - if expression.pattern() in [ - r'"[^"\\]*(\\.[^"\\]*)*"', - r"'[^'\\]*(\\.[^'\\]*)*'", - ]: - innerIndex = self.tri_single[0].indexIn(text, index + 1) - if innerIndex == -1: - innerIndex = self.tri_double[0].indexIn(text, index + 1) - - if innerIndex != -1: - tripleQuoteIndexes = range(innerIndex, innerIndex + 3) - self.tripleQuoutesWithinStrings.extend(tripleQuoteIndexes) + # if there is a string we check + # if there are some triple quotes within the string + # they will be ignored if they are matched again + if index >= 0 and expression.pattern() in [ + r'"[^"\\]*(\\.[^"\\]*)*"', + r"'[^'\\]*(\\.[^'\\]*)*'", + ]: + innerIndex = self.tri_single[0].indexIn(text, index + 1) + if innerIndex == -1: + innerIndex = self.tri_double[0].indexIn(text, index + 1) + + if innerIndex != -1: + tripleQuoteIndexes = range(innerIndex, innerIndex + 3) + self.tripleQuoutesWithinStrings.extend(tripleQuoteIndexes) while index >= 0: # skipping triple quotes within strings @@ -237,7 +236,4 @@ def match_multiline(self, text, delimiter, in_state, style): start = delimiter.indexIn(text, start + length) # Return True if still inside a multi-line string, False otherwise - if self.currentBlockState() == in_state: - return True - else: - return False + return self.currentBlockState() == in_state diff --git a/src/badger/gui/components/var_table.py b/src/badger/gui/components/var_table.py index b9b5114c..67247205 100644 --- a/src/badger/gui/components/var_table.py +++ b/src/badger/gui/components/var_table.py @@ -21,44 +21,43 @@ checked variables with their current bounds. """ +import logging +import traceback from functools import partial from importlib import resources -import traceback from typing import Any, cast + +from gest_api.vocs import ContinuousVariable +from PyQt5.QtCore import QPoint, QSize, Qt, pyqtSignal +from PyQt5.QtGui import ( + QColor, + QDragEnterEvent, + QDragMoveEvent, + QDropEvent, + QGuiApplication, + QIcon, +) from PyQt5.QtWidgets import ( - QTableWidget, - QTableWidgetItem, - QHeaderView, + QAbstractItemView, QCheckBox, + QDialog, + QGridLayout, + QHBoxLayout, + QHeaderView, + QLabel, + QMenu, QMessageBox, - QAbstractItemView, QPushButton, + QTableWidget, + QTableWidgetItem, QWidget, - QHBoxLayout, - QMenu, - QGridLayout, - QLabel, - QDialog, -) -from PyQt5.QtCore import pyqtSignal, Qt, QSize, QPoint -from PyQt5.QtGui import ( - QColor, - QIcon, - QGuiApplication, - QDropEvent, - QDragMoveEvent, - QDragEnterEvent, ) -from badger.gui.components.robust_spinbox import RobustSpinBox from badger.environment import Environment, instantiate_env from badger.errors import BadgerInterfaceChannelError +from badger.gui.components.robust_spinbox import RobustSpinBox from badger.gui.windows.expandable_message_box import ExpandableMessageBox -from gest_api.vocs import ContinuousVariable - -import logging - logger = logging.getLogger(__name__) @@ -138,7 +137,7 @@ def update_vocs(self): """ Emit the data_changed signal to notify that the VOCS has been updated. """ - logging.debug("Emitting data_changed signal from VariableTable") + logger.debug("Emitting data_changed signal from VariableTable") self.data_changed.emit() def config_logic(self): @@ -160,7 +159,7 @@ def is_all_checked(self): for i in range(self.rowCount() - 1): item = self.cellWidget(i, 0) if not item: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") item = cast(QCheckBox, item) if not item.isChecked(): return False @@ -176,7 +175,7 @@ def header_clicked(self, idx: int): for i in range(self.rowCount() - 1): item = self.cellWidget(i, 0) if not item: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") item = cast(QCheckBox, item) # Doing batch update item.blockSignals(True) @@ -188,15 +187,15 @@ def update_bounds(self): for i in range(self.rowCount() - 1): widget = self.item(i, 1) if not widget: - raise Exception("Variable name widget not found!") + raise RuntimeError("Variable name widget not found!") name = widget.text() sb_lower = self.cellWidget(i, 2) if not sb_lower: - raise Exception("Variable bound spinbox widget not found!") + raise RuntimeError("Variable bound spinbox widget not found!") sb_upper = self.cellWidget(i, 3) if not sb_upper: - raise Exception("Variable bound spinbox widget not found!") + raise RuntimeError("Variable bound spinbox widget not found!") sb_lower = cast(RobustSpinBox, sb_lower) sb_upper = cast(RobustSpinBox, sb_upper) @@ -211,10 +210,10 @@ def validate_row(self, row: int): """ sb_lower = self.cellWidget(row, 2) # Min value spinbox if not sb_lower: - raise Exception("Variable bound spinbox widget not found!") + raise RuntimeError("Variable bound spinbox widget not found!") sb_upper = self.cellWidget(row, 3) # Max value spinbox if not sb_upper: - raise Exception("Variable bound spinbox widget not found!") + raise RuntimeError("Variable bound spinbox widget not found!") sb_lower = cast(RobustSpinBox, sb_lower) sb_upper = cast(RobustSpinBox, sb_upper) @@ -228,8 +227,8 @@ def validate_row(self, row: int): def set_bounds( self, variables: dict[str, tuple[float, float]], signal: bool = True ): - for name in variables: - self.bounds[name] = variables[name] + for name, bounds in variables.items(): + self.bounds[name] = bounds if signal: self.update_variables(self.variables, 2) @@ -259,12 +258,12 @@ def update_selected(self): for i in range(self.rowCount() - 1): _cb = self.cellWidget(i, 0) if not _cb: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") _cb = cast(QCheckBox, _cb) widget = self.item(i, 1) if not widget: - raise Exception("Variable name widget not found!") + raise RuntimeError("Variable name widget not found!") name = widget.text() is_selected = _cb.isChecked() widget.setForeground(QColor("lightgray" if is_selected else "gray")) @@ -395,7 +394,7 @@ def update_variables( _cb = self.cellWidget(i, 0) if not _cb: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") _cb = cast(QCheckBox, _cb) _cb.setChecked(self.is_checked(name)) @@ -525,7 +524,7 @@ def try_insert_variable(self, name: str): f"Variable {name} cannot be found through the interface!", ) return - except Exception: + except Exception: # noqa: BLE001 - interface bounds fetch varies # Raised when PV exists but value/hard limits cannot be found # Set to some default values _bounds = [0, 0] @@ -542,14 +541,14 @@ def try_insert_variable(self, name: str): else: # TODO: handle this case? Right now I don't think it should happen - raise Exception("Environment cannot be found for new variable bounds!") + raise RuntimeError("Environment cannot be found for new variable bounds!") # Add checkbox only when a PV is entered self.setCellWidget(idx, 0, QCheckBox()) _cb = self.cellWidget(idx, 0) if not _cb: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") _cb = cast(QCheckBox, _cb) # Checked by default when entered diff --git a/src/badger/gui/mini/__init__.py b/src/badger/gui/mini/__init__.py index c8d508fb..797b633a 100644 --- a/src/badger/gui/mini/__init__.py +++ b/src/badger/gui/mini/__init__.py @@ -1,23 +1,23 @@ """Entry points and app bootstrap helpers for the Badger mini GUI package.""" -from importlib import resources -import signal +import logging import os +import signal import sys import time -from PyQt5.QtWidgets import QApplication -from PyQt5.QtGui import QFont, QIcon -from PyQt5 import QtCore -from qdarkstyle import load_stylesheet, DarkPalette +import traceback +from importlib import resources +from types import TracebackType +from typing import NoReturn -from badger.settings import init_settings -from badger.gui.mini.windows.main_window import BadgerMiniWindow +from PyQt5 import QtCore +from PyQt5.QtGui import QFont, QIcon +from PyQt5.QtWidgets import QApplication +from qdarkstyle import DarkPalette, load_stylesheet -import traceback from badger.errors import BadgerError -from types import TracebackType -from typing import Type, NoReturn -import logging +from badger.gui.mini.windows.main_window import BadgerMiniWindow +from badger.settings import init_settings logger = logging.getLogger(__name__) @@ -52,7 +52,7 @@ def on_timeout(): def error_handler( - etype: Type[BaseException], value: BaseException, tb: TracebackType + etype: type[BaseException], value: BaseException, tb: TracebackType ) -> NoReturn: """ Custom exception handler that formats uncaught exceptions and raises a BadgerError. @@ -76,7 +76,7 @@ def error_handler( raise BadgerError(error_title, error_msg) -def launch_gui(config_path=None, template_filename: str = None): +def launch_gui(config_path=None, template_filename: str | None = None): sys.excepthook = error_handler app = QApplication(sys.argv) diff --git a/src/badger/gui/mini/components/env_cbox.py b/src/badger/gui/mini/components/env_cbox.py index 8074ba77..dc19cecf 100644 --- a/src/badger/gui/mini/components/env_cbox.py +++ b/src/badger/gui/mini/components/env_cbox.py @@ -7,46 +7,42 @@ vocs tables emit ``vocs_updated`` so the rest of the UI stays in sync. """ +import logging from pathlib import Path from typing import Any +import numpy as np +from gest_api.vocs import VOCS, ContinuousVariable +from pydantic_core import ValidationError +from PyQt5.QtCore import QRegExp, pyqtSignal from PyQt5.QtWidgets import ( + QCheckBox, QFrame, - QVBoxLayout, QHBoxLayout, + QLabel, + QLineEdit, QPushButton, QStyle, + QStyledItemDelegate, QStyleOptionComboBox, - QWidget, - QLineEdit, QTreeWidget, + QVBoxLayout, + QWidget, ) -from PyQt5.QtWidgets import ( - QCheckBox, - QStyledItemDelegate, - QLabel, -) -from PyQt5.QtCore import QRegExp, pyqtSignal -import numpy as np -from badger.gui.components.obs_table import ObservableTable -from badger.settings import init_settings -from pydantic_core import ValidationError from badger.errors import BadgerRoutineError from badger.gui.components.collapsible_box import CollapsibleBox -from badger.gui.components.pydantic_editor import BadgerPydanticEditor -from badger.gui.mini.components.var_table import VariableTable -from badger.gui.components.obj_table import ObjectiveTable from badger.gui.components.con_table import ConstraintTable from badger.gui.components.data_table import init_data_table +from badger.gui.components.obj_table import ObjectiveTable +from badger.gui.components.obs_table import ObservableTable +from badger.gui.components.pydantic_editor import BadgerPydanticEditor +from badger.gui.mini.components.var_table import VariableTable from badger.gui.utils import ( MouseWheelWidgetAdjustmentGuard, NoHoverFocusComboBox, ) -from xopt.vocs import VOCS -from gest_api.vocs import ContinuousVariable - -import logging +from badger.settings import init_settings LABEL_WIDTH = 96 @@ -66,8 +62,7 @@ class ArrowOnlyPopupComboBox(NoHoverFocusComboBox): def __init__(self, parent=None): super().__init__(parent) - self.setStyleSheet( - """ + self.setStyleSheet(""" QComboBox { color: darkGray; background-color: transparent; @@ -78,8 +73,7 @@ def __init__(self, parent=None): border: none; width: 14px; } - """ - ) + """) self.setItemDelegate(QStyledItemDelegate()) self.installEventFilter(MouseWheelWidgetAdjustmentGuard(self)) @@ -232,9 +226,14 @@ class BadgerEnvBox(QWidget): def __init__( self, parent: QWidget | None = None, - envs: list[str] = [], - generators: list[str] = [], + envs: list[str] | None = None, + generators: list[str] | None = None, ): + if envs is None: + envs = [] + if generators is None: + generators = [] + super().__init__(parent) self.envs = envs diff --git a/src/badger/gui/mini/components/var_table.py b/src/badger/gui/mini/components/var_table.py index f4cc9a23..af885087 100644 --- a/src/badger/gui/mini/components/var_table.py +++ b/src/badger/gui/mini/components/var_table.py @@ -8,36 +8,39 @@ RangeAdjustButtonStack helper. """ -from functools import partial -from importlib import resources +import logging import math import traceback +from functools import partial +from importlib import resources from typing import Any, cast + +from gest_api.vocs import ContinuousVariable +from PyQt5.QtCore import QPoint, QSize, Qt, QTimer, pyqtSignal +from PyQt5.QtGui import ( + QColor, + QDragEnterEvent, + QDragMoveEvent, + QDropEvent, + QGuiApplication, + QIcon, +) from PyQt5.QtWidgets import ( - QTableWidget, - QTableWidgetItem, - QHeaderView, - QCheckBox, - QMessageBox, QAbstractItemView, - QPushButton, - QWidget, + QCheckBox, + QDialog, + QGridLayout, QHBoxLayout, - QVBoxLayout, + QHeaderView, + QLabel, QLineEdit, QMenu, - QGridLayout, - QLabel, - QDialog, -) -from PyQt5.QtCore import pyqtSignal, Qt, QSize, QPoint, QTimer -from PyQt5.QtGui import ( - QColor, - QIcon, - QGuiApplication, - QDropEvent, - QDragMoveEvent, - QDragEnterEvent, + QMessageBox, + QPushButton, + QTableWidget, + QTableWidgetItem, + QVBoxLayout, + QWidget, ) from badger.environment import Environment, instantiate_env @@ -45,10 +48,6 @@ from badger.gui.windows.expandable_message_box import ExpandableMessageBox from badger.utils import _round_bounds_inward -from gest_api.vocs import ContinuousVariable - -import logging - logger = logging.getLogger(__name__) @@ -172,8 +171,11 @@ def __init__( on_increase=None, on_decrease=None, parent: QWidget | None = None, - button_size: QSize = QSize(12, 10), + button_size: QSize | None = None, ): + if button_size is None: + button_size = QSize(12, 10) + super().__init__(parent) self.up_btn = QPushButton("▲") @@ -443,7 +445,7 @@ def update_vocs(self): """ Emit the data_changed signal to notify that the VOCS has been updated. """ - logging.debug("Emitting data_changed signal from VariableTable") + logger.debug("Emitting data_changed signal from VariableTable") self.data_changed.emit() def config_logic(self): @@ -518,7 +520,7 @@ def is_all_checked(self): for i in range(self.rowCount() - 1): item = self.cellWidget(i, 0) if not item: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") item = cast(QCheckBox, item) if not item.isChecked(): return False @@ -534,7 +536,7 @@ def header_clicked(self, idx: int): for i in range(self.rowCount() - 1): item = self.cellWidget(i, 0) if not item: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") item = cast(QCheckBox, item) # Doing batch update item.blockSignals(True) @@ -557,10 +559,10 @@ def set_bounds( variables: dict[str, tuple[float, float]], signal: bool = True, clipped: dict[str, bool] | None = None, - ): + ) -> None: clipped = clipped or {} - for name in variables: - self.bounds[name] = variables[name] + for name, value in variables.items(): + self.bounds[name] = value if name in clipped: self.clipped[name] = clipped[name] else: @@ -599,12 +601,12 @@ def update_selected(self): for i in range(self.rowCount() - 1): _cb = self.cellWidget(i, 0) if not _cb: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") _cb = cast(QCheckBox, _cb) widget = self.item(i, 1) if not widget: - raise Exception("Variable name widget not found!") + raise RuntimeError("Variable name widget not found!") name = widget.text() is_selected = _cb.isChecked() widget.setForeground(QColor("lightgray" if is_selected else "gray")) @@ -834,7 +836,7 @@ def update_variables( _cb = self.cellWidget(i, 0) if not _cb: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") _cb = cast(QCheckBox, _cb) _cb.setChecked(self.is_checked(name)) @@ -960,7 +962,7 @@ def try_insert_variable(self, name: str): f"Variable {name} cannot be found through the interface!", ) return - except Exception: + except Exception: # noqa: BLE001 - interface bounds fetch varies # Raised when PV exists but value/hard limits cannot be found # Set to some default values _bounds = [0, 0] @@ -978,14 +980,14 @@ def try_insert_variable(self, name: str): else: # TODO: handle this case? Right now I don't think it should happen - raise Exception("Environment cannot be found for new variable bounds!") + raise RuntimeError("Environment cannot be found for new variable bounds!") # Add checkbox only when a PV is entered self.setCellWidget(idx, 0, QCheckBox()) _cb = self.cellWidget(idx, 0) if not _cb: - raise Exception("Checkbox widget not found!") + raise RuntimeError("Checkbox widget not found!") _cb = cast(QCheckBox, _cb) # Checked by default when entered diff --git a/src/badger/gui/mini/pages/home_page.py b/src/badger/gui/mini/pages/home_page.py index 151e3856..350b9d06 100644 --- a/src/badger/gui/mini/pages/home_page.py +++ b/src/badger/gui/mini/pages/home_page.py @@ -7,30 +7,32 @@ """ import gc +import logging import os import traceback from importlib import resources import numpy as np from pandas import DataFrame -from PyQt5.QtCore import pyqtSignal, Qt, QModelIndex +from PyQt5.QtCore import QModelIndex, Qt, pyqtSignal from PyQt5.QtGui import QIcon from PyQt5.QtWidgets import ( + QLabel, QMessageBox, QSplitter, QVBoxLayout, QWidget, - QLabel, ) from badger.archive import ( delete_run, get_base_run_filename, - load_run, get_runs, + load_run, save_tmp_run, ) from badger.errors import BadgerRoutineError +from badger.gui.components.action_bar import BadgerActionBar from badger.gui.components.data_panel import filter_metadata from badger.gui.components.data_table import ( add_row, @@ -38,23 +40,19 @@ reset_table, update_table, ) -from badger.gui.mini.pages.routine_page import BadgerRoutinePage - from badger.gui.components.navigators import TemplateNavigator from badger.gui.components.run_monitor import BadgerOptMonitor from badger.gui.components.status_bar import BadgerStatusBar -from badger.gui.components.action_bar import BadgerActionBar -from badger.utils import get_header -from badger.settings import init_settings +from badger.gui.mini.pages.routine_page import BadgerRoutinePage +from badger.gui.utils import build_bax_results_file # from PyQt5.QtGui import QBrush, QColor from badger.gui.windows.message_dialog import BadgerScrollableMessageBox from badger.gui.windows.terminition_condition_dialog import ( BadgerTerminationConditionDialog, ) -from badger.gui.utils import ModalOverlay, build_bax_results_file - -import logging +from badger.settings import init_settings +from badger.utils import get_header logger = logging.getLogger(__name__) @@ -146,13 +144,11 @@ def init_ui(self): vbox_table = QVBoxLayout(panel_table) vbox_table.setContentsMargins(0, 0, 0, 0) title_label = QLabel("Run Data") - title_label.setStyleSheet( - """ + title_label.setStyleSheet(""" background-color: #455364; font-weight: bold; padding: 4px; - """ - ) + """) title_label.setAlignment(Qt.AlignCenter) # Center-align the title vbox_table.addWidget(title_label, 0) self.run_table = run_table = data_table() @@ -329,7 +325,7 @@ def init_home_page(self): self.routine_editor.env_box.algo_cb.setCurrentIndex(idx) self.routine_editor.select_generator(idx) - def go_run(self, i: int = None): + def go_run(self, i: int | None = None): logger.info(f"Activating run: {i}") gc.collect() @@ -358,7 +354,7 @@ def go_run(self, i: int = None): self.run_monitor.routine_filename = run_filename except IndexError: return - except Exception as e: # failed to load the run + except Exception as e: # noqa: BLE001 - run load boundary details = traceback.format_exc() dialog = BadgerScrollableMessageBox( title="Error!", text=str(e), parent=self @@ -509,9 +505,9 @@ def prepare_run(self, data=None, init_points_flag=True): logger.info("Preparing new run.") try: routine = self.routine_editor._compose_routine() - except Exception as e: + except Exception: self.sig_routine_invalid.emit() - raise e + raise # Give this run its own results folder, named after the run's archive # name (-) so the folder used during the run matches @@ -707,18 +703,17 @@ def delete_run(self): def cover_page(self): logger.info("Covering page with overlay.") - return # disable overlay for now - try: - self.overlay - except AttributeError: - # Set parent to the main window - try: - main_window = self.parent().parent() - except AttributeError: # in test mode - return - self.overlay = ModalOverlay(main_window) - self.overlay.show() + # try: + # self.overlay + # except AttributeError: + # # Set parent to the main window + # try: + # main_window = self.parent().parent() + # except AttributeError: # in test mode + # return + # self.overlay = ModalOverlay(main_window) + # self.overlay.show() def uncover_page(self): logger.info("Uncovering page overlay.") diff --git a/src/badger/gui/mini/pages/routine_page.py b/src/badger/gui/mini/pages/routine_page.py index 9216c52c..ebf52b69 100644 --- a/src/badger/gui/mini/pages/routine_page.py +++ b/src/badger/gui/mini/pages/routine_page.py @@ -8,83 +8,87 @@ an existing Routine back into the form. """ -from typing import Any -import warnings -import traceback import copy -from functools import partial +import logging import os -import yaml +import traceback +import warnings +from datetime import UTC, datetime +from functools import partial +from typing import Any import numpy as np import pandas as pd -from PyQt5.QtCore import pyqtSignal, QTimer -from PyQt5.QtWidgets import QLineEdit, QPushButton, QFileDialog -from PyQt5.QtWidgets import QMessageBox, QWidget, QTabWidget -from PyQt5.QtWidgets import QVBoxLayout, QScrollArea -from PyQt5.QtWidgets import QTableWidgetItem, QPlainTextEdit -from PyQt5.QtWidgets import QApplication -from badger.gui.components.navigators import HistoryNavigator +import yaml from coolname import generate_slug -from xopt import VOCS -from xopt.vocs import random_inputs -from xopt.generators import ( - get_generator_defaults, - all_generator_names, - get_generator_dynamic, -) -from xopt.vocs import get_local_region from gest_api.vocs import ( BaseObjective, + ContinuousVariable, GreaterThanConstraint, LessThanConstraint, MaximizeObjective, MinimizeObjective, - ContinuousVariable, ) - from pydantic import ValidationError +from PyQt5.QtCore import QTimer, pyqtSignal +from PyQt5.QtWidgets import ( + QApplication, + QFileDialog, + QLineEdit, + QMessageBox, + QPlainTextEdit, + QPushButton, + QScrollArea, + QTableWidgetItem, + QTabWidget, + QVBoxLayout, + QWidget, +) +from xopt import VOCS +from xopt.generators import ( + all_generator_names, + get_generator_defaults, + get_generator_dynamic, +) +from xopt.vocs import get_local_region, random_inputs +from badger.environment import instantiate_env +from badger.errors import ( + BadgerEnvInstantiationError, + BadgerEnvNotFoundError, + BadgerEnvVarError, + BadgerRoutineError, + VariableRangeError, +) +from badger.factory import get_env, list_env, list_generators from badger.gui.components.data_panel import BadgerDataPanel from badger.gui.components.data_table import ( get_table_content_as_dict, set_init_data_table, update_init_data_table, ) +from badger.gui.components.navigators import HistoryNavigator from badger.gui.mini.components.env_cbox import BadgerEnvBox +from badger.gui.utils import filter_generator_config, with_busy_cursor +from badger.gui.windows.add_random_dialog import BadgerAddRandomDialog from badger.gui.windows.docs_window import BadgerDocsWindow -from badger.gui.windows.lim_vrange_dialog import BadgerLimitVariableRangeDialog from badger.gui.windows.ind_lim_vrange_dialog import ( BadgerIndividualLimitVariableRangeDialog, ) -from badger.gui.windows.review_dialog import BadgerReviewDialog -from badger.gui.windows.add_random_dialog import BadgerAddRandomDialog +from badger.gui.windows.lim_vrange_dialog import BadgerLimitVariableRangeDialog from badger.gui.windows.message_dialog import BadgerScrollableMessageBox -from badger.gui.utils import filter_generator_config, with_busy_cursor -from badger.environment import instantiate_env -from badger.errors import ( - BadgerEnvNotFoundError, - BadgerRoutineError, - BadgerEnvVarError, - BadgerEnvInstantiationError, - VariableRangeError, -) -from badger.factory import list_generators, list_env, get_env +from badger.gui.windows.review_dialog import BadgerReviewDialog from badger.routine import Routine from badger.settings import init_settings -from datetime import datetime from badger.utils import ( BlockSignalsContext, - load_config, + _round_bounds_inward, get_badger_version, get_xopt_version, + load_config, ts_float_to_str, - _round_bounds_inward, ) - -import logging - logger = logging.getLogger(__name__) @@ -119,7 +123,7 @@ def extract_objective_symbol(objective: BaseObjective) -> str: if isinstance(objective, MaximizeObjective): return "MAXIMIZE" else: - raise ValueError(f"Unknown objective type: {objective}") + raise TypeError(f"Unknown objective type: {objective}") class BadgerRoutinePage(QWidget): @@ -203,8 +207,7 @@ def init_ui(self): self.env_box = BadgerEnvBox(None, self.envs, self.generators) scroll_area = QScrollArea() scroll_area.setFrameShape(QScrollArea.NoFrame) - scroll_area.setStyleSheet( - """ + scroll_area.setStyleSheet(""" QScrollArea { border: none; /* Remove border */ margin: 0px; /* Remove margin */ @@ -213,8 +216,7 @@ def init_ui(self): QScrollArea > QWidget { margin: 0px; /* Remove margin inside */ } - """ - ) + """) scroll_content_env = QWidget() scroll_layout_env = QVBoxLayout(scroll_content_env) # add extra right margin for macOS to prevent scrollbar overlap @@ -642,28 +644,27 @@ def _filter_generator_params(self, generator_name: str, generator_config: dict): Filter which generator parameters get saved to template """ - if generator_name in ["expected_improvement", "upper_confidence_bound"]: - if ( - "turbo_controller" in generator_config - and generator_config["turbo_controller"] is not None - and isinstance(generator_config["turbo_controller"], dict) - ): - turbo = generator_config["turbo_controller"] - generator_config["turbo_controller"] = { - k: v - for k, v in turbo.items() - if k - in { - "name", - "length", - "length_max", - "length_min", - "failure_tolerance", - "success_tolerance", - "scale_factor", - "restrict_model_data", - } + if generator_name in ["expected_improvement", "upper_confidence_bound"] and ( + "turbo_controller" in generator_config + and generator_config["turbo_controller"] is not None + and isinstance(generator_config["turbo_controller"], dict) + ): + turbo = generator_config["turbo_controller"] + generator_config["turbo_controller"] = { + k: v + for k, v in turbo.items() + if k + in { + "name", + "length", + "length_max", + "length_min", + "failure_tolerance", + "success_tolerance", + "scale_factor", + "restrict_model_data", } + } return generator_config @@ -679,7 +680,7 @@ def save_template_yaml(self): # Suggest a filename based on the routine name or placeholder routine_name = self.edit_save.text() or self.edit_save.placeholderText() if not routine_name: - routine_name = "template_" + datetime.now().strftime("%y%m%d_%H%M%S") + routine_name = "template_" + datetime.now(tz=UTC).strftime("%y%m%d_%H%M%S") suggested_filename = f"{routine_name}.yaml" template_path, _ = QFileDialog.getSaveFileName( self, @@ -999,7 +1000,7 @@ def select_generator(self, i: int): # Get vocs try: vocs, _ = self.env_box.compose_vocs() - except Exception: + except BadgerRoutineError: vocs = None self.env_box.edit_algo_params.set_params_from_generator( name, filtered_config, vocs @@ -1037,7 +1038,9 @@ def create_env(self): try: env = instantiate_env(self.env, configs) except Exception as e: - raise BadgerEnvInstantiationError(f"Failed to instantiate environment: {e}") + raise BadgerEnvInstantiationError( + f"Failed to instantiate environment: {e}" + ) from e return env @@ -1047,10 +1050,11 @@ def refresh_params_generator(self): try: tmp = {} - exec(self.script, tmp) + # User-provided script must define a `generate` function, so exec is required here. + exec(self.script, tmp) # noqa: S102 try: tmp["generate"] # test if generate function is defined - except Exception as e: + except KeyError as e: QMessageBox.warning( self, "Please define a valid generate function!", str(e) ) @@ -1060,14 +1064,14 @@ def refresh_params_generator(self): # Get vocs try: vocs, _ = self.env_box.compose_vocs() - except Exception: + except BadgerRoutineError: vocs = None # Function generate comes from the script params_generator = tmp["generate"](env, vocs) self.env_box.edit_algo_params.set_params_from_generator( self.routine.generator.name, params_generator, vocs ) - except Exception as e: + except Exception as e: # noqa: BLE001 - runs user-provided generator script QMessageBox.warning(self, "Invalid script!", str(e)) @with_busy_cursor @@ -1105,7 +1109,7 @@ def select_env(self, i: int): self.env = env self.env_box.edit_var.clear() self.env_box.edit_obj.clear() - except Exception: + except Exception: # noqa: BLE001 - env selection rollback on any failure self.configs = None self.env = None self.env_box.clear_selected_env() @@ -1238,7 +1242,7 @@ def fill_curr_in_init_table(self, record=False): raise BadgerEnvVarError( f"Failed to get current variable values : {e}\n" "Please ensure the environment is properly configured." - ) + ) from e # Iterate through the rows for row in range(table.rowCount()): @@ -1422,7 +1426,7 @@ def set_vrange(self, set_all=True): raise BadgerEnvVarError( f"Failed to get current variable values : {e}\n" "Please ensure the environment is properly configured." - ) + ) from e option_idx = self.limit_option["limit_option_idx"] clipped = {} @@ -1463,7 +1467,7 @@ def set_vrange(self, set_all=True): self.update_init_table() # auto populate if option is set # remember user selection for applying limit changes - if not self.lim_apply_to_vars == 2: + if self.lim_apply_to_vars != 2: # Check if lim_apply_to_vars has been initialized # It will be set to 2 until the btn_lim_vrange is clicked self.lim_apply_to_vars = set_all @@ -1771,7 +1775,7 @@ def toggle_relative_to_curr(self, checked, refresh=True): if checked: try: _ = self.env_box.compose_vocs() - except Exception: + except BadgerRoutineError: logger.warning("Variable range is not valid, switching to manual mode.") QTimer.singleShot( 0, lambda: self.env_box.relative_to_curr.isChecked() @@ -1943,10 +1947,9 @@ def _compose_routine(self) -> Routine: NO_OBJECTIVE_GENERATORS = ["bax"] - if not vocs.objectives: - if generator_name not in NO_OBJECTIVE_GENERATORS: - logger.error("No objectives selected.") - raise BadgerRoutineError("no objectives selected") + if not vocs.objectives and generator_name not in NO_OBJECTIVE_GENERATORS: + logger.error("No objectives selected.") + raise BadgerRoutineError("no objectives selected") # Initial points init_points_df = pd.DataFrame.from_dict( @@ -2005,7 +2008,9 @@ def _compose_routine(self) -> Routine: # Metadata badger_version=get_badger_version(), xopt_version=get_xopt_version(), - creation_ts=ts_float_to_str(datetime.now().timestamp(), "lcls-fname"), + creation_ts=ts_float_to_str( + datetime.now(tz=UTC).timestamp(), "lcls-fname" + ), # Xopt part generator=generator, # Badger part @@ -2044,7 +2049,7 @@ def _compose_routine(self) -> Routine: def review(self): try: routine = self._compose_routine() - except Exception: + except Exception: # noqa: BLE001 - routine compose reports via dialog return QMessageBox.critical( self, "Invalid routine!", traceback.format_exc() ) diff --git a/src/badger/gui/mini/windows/main_window.py b/src/badger/gui/mini/windows/main_window.py index 1b3c4fec..6f338945 100644 --- a/src/badger/gui/mini/windows/main_window.py +++ b/src/badger/gui/mini/windows/main_window.py @@ -2,9 +2,10 @@ import logging from importlib import metadata -from typing import Dict + from PyQt5.QtCore import QThread from PyQt5.QtWidgets import QDesktopWidget, QMainWindow, QMessageBox, QStackedWidget + from badger.gui.components.create_process import CreateProcess from badger.gui.components.process_manager import ProcessManager from badger.gui.mini.pages.home_page import BadgerHomePage @@ -55,7 +56,7 @@ def cleanupThread(self) -> None: if thread in self.thread_list: self.thread_list.remove(thread) - def storeSubprocess(self, process_with_args: Dict) -> None: + def storeSubprocess(self, process_with_args: dict) -> None: logger.info(f"Storing prepared subprocess: {process_with_args}") """ Store the prepared subprocess for later use. diff --git a/src/badger/gui/pages/home_page.py b/src/badger/gui/pages/home_page.py index f0603340..8cda47ee 100644 --- a/src/badger/gui/pages/home_page.py +++ b/src/badger/gui/pages/home_page.py @@ -51,7 +51,7 @@ from badger.gui.components.routine_page import BadgerRoutinePage from badger.gui.components.run_monitor import BadgerOptMonitor from badger.gui.components.status_bar import BadgerStatusBar -from badger.gui.utils import ModalOverlay, build_bax_results_file +from badger.gui.utils import build_bax_results_file # from PyQt5.QtGui import QBrush, QColor from badger.gui.windows.message_dialog import BadgerScrollableMessageBox @@ -88,7 +88,7 @@ class BadgerHomePage(QWidget): sig_routine_activated = pyqtSignal(bool) sig_routine_invalid = pyqtSignal() - def __init__(self, process_manager: "Optional[ProcessManager]" = None): + def __init__(self, process_manager: "ProcessManager | None" = None): logger.info("Initializing BadgerHomePage.") super().__init__() @@ -312,7 +312,7 @@ def init_home_page(self) -> None: # Load the default generator self.routine_editor.set_default_generator("neldermead") - def go_run(self, i: int = None) -> None: + def go_run(self, i: int | None = None) -> None: logger.info(f"Activating run: {i}") gc.collect() @@ -342,7 +342,7 @@ def go_run(self, i: int = None) -> None: self.run_monitor.routine_filename = run_filename except IndexError: return - except Exception as e: # failed to load the run + except Exception as e: # noqa: BLE001 - run load boundary details = traceback.format_exc() dialog = BadgerScrollableMessageBox( title="Error!", text=str(e), parent=self @@ -496,9 +496,9 @@ def prepare_run( logger.info("Preparing new run.") try: routine = self.routine_editor._compose_routine() - except Exception as e: + except Exception: self.sig_routine_invalid.emit() - raise e + raise # Give this run its own results folder, named after the run's archive # name (-) so the folder used during the run matches @@ -660,18 +660,17 @@ def delete_run(self) -> None: def cover_page(self) -> None: logger.info("Covering page with overlay.") - return # disable overlay for now - try: - self.overlay - except AttributeError: - # Set parent to the main window - try: - main_window = self.parent().parent() - except AttributeError: # in test mode - return - self.overlay = ModalOverlay(main_window) - self.overlay.show() + # try: + # self.overlay + # except AttributeError: + # # Set parent to the main window + # try: + # main_window = self.parent().parent() + # except AttributeError: # in test mode + # return + # self.overlay = ModalOverlay(main_window) + # self.overlay.show() def uncover_page(self) -> None: logger.info("Uncovering page overlay.") diff --git a/src/badger/gui/utils.py b/src/badger/gui/utils.py index afe66894..d9e276ef 100644 --- a/src/badger/gui/utils.py +++ b/src/badger/gui/utils.py @@ -5,21 +5,22 @@ import copy import logging import os -from importlib import resources -from typing import Any, Callable +from collections.abc import Callable from functools import wraps +from importlib import resources +from typing import Any from PyQt5.QtCore import QEvent, QObject, QSize, Qt from PyQt5.QtGui import QIcon from PyQt5.QtWidgets import ( QAbstractSpinBox, + QApplication, QComboBox, QDialog, QLabel, QPushButton, QToolButton, QVBoxLayout, - QApplication, ) from badger.errors import BadgerConfigError diff --git a/src/badger/gui/windows/add_random_dialog.py b/src/badger/gui/windows/add_random_dialog.py index ee19674f..fcc73bb2 100644 --- a/src/badger/gui/windows/add_random_dialog.py +++ b/src/badger/gui/windows/add_random_dialog.py @@ -4,10 +4,21 @@ from copy import deepcopy -from PyQt5.QtWidgets import QDialog, QWidget, QHBoxLayout, QStackedWidget -from PyQt5.QtWidgets import QVBoxLayout, QSpinBox, QPushButton -from PyQt5.QtWidgets import QGroupBox, QLabel, QComboBox, QStyledItemDelegate from PyQt5.QtCore import Qt +from PyQt5.QtWidgets import ( + QComboBox, + QDialog, + QGroupBox, + QHBoxLayout, + QLabel, + QPushButton, + QSpinBox, + QStackedWidget, + QStyledItemDelegate, + QVBoxLayout, + QWidget, +) + from badger.gui.components.robust_spinbox import RobustSpinBox diff --git a/src/badger/gui/windows/docs_window.py b/src/badger/gui/windows/docs_window.py index 94edf1da..e3d1c614 100644 --- a/src/badger/gui/windows/docs_window.py +++ b/src/badger/gui/windows/docs_window.py @@ -2,16 +2,18 @@ environments, and general Badger guides in a QTextBrowser with clickable navigation links.""" +from typing import TYPE_CHECKING + from PyQt5.QtWidgets import ( - QHBoxLayout, - QVBoxLayout, QCheckBox, - QWidget, + QHBoxLayout, QMainWindow, QTextBrowser, + QVBoxLayout, + QWidget, ) -from badger.factory import load_badger_docs, load_plugin_docs, list_generators -from typing import TYPE_CHECKING + +from badger.factory import list_generators, load_badger_docs, load_plugin_docs if TYPE_CHECKING: from PyQt5.QtCore import QUrl @@ -72,7 +74,7 @@ def config_logic(self): self.cb_md.stateChanged.connect(self.refresh_docs_view) self.markdown_viewer.anchorClicked.connect(self.handle_link_click) - def load_docs(self, subdir: str = None): + def load_docs(self, subdir: str | None = None): """ Load the docs for the current generator and subdir (if provided). @@ -88,7 +90,7 @@ def load_docs(self, subdir: str = None): else: # Load plugin documentation from plugin root self.docs = load_plugin_docs(self.docs_name, self.plugin_type) - except Exception as e: + except Exception as e: # noqa: BLE001 - docs viewer shows any load error self.docs = str(e) self.refresh_docs_view() @@ -122,7 +124,7 @@ def handle_link_click(self, url: "QUrl") -> None: href = url.toString() # Indicate links not yet supported in GUI docs viewer - if href.startswith("https://") or href.startswith("mailto:"): + if href.startswith(("https://", "mailto:")): self.docs_name = "external links not yet implemented in GUI docs viewer" self.load_docs() return diff --git a/src/badger/gui/windows/edit_script_dialog.py b/src/badger/gui/windows/edit_script_dialog.py index c18de9a7..32c5f48c 100644 --- a/src/badger/gui/windows/edit_script_dialog.py +++ b/src/badger/gui/windows/edit_script_dialog.py @@ -2,9 +2,16 @@ Python code editor for writing custom generator scripts that execute as part of an optimization routine.""" -from PyQt5.QtWidgets import QDialog, QPlainTextEdit, QVBoxLayout, QWidget -from PyQt5.QtWidgets import QHBoxLayout, QPushButton from PyQt5.QtGui import QFont +from PyQt5.QtWidgets import ( + QDialog, + QHBoxLayout, + QPlainTextEdit, + QPushButton, + QVBoxLayout, + QWidget, +) + from badger.gui.components.syntax import PythonHighlighter diff --git a/src/badger/gui/windows/expandable_message_box.py b/src/badger/gui/windows/expandable_message_box.py index 52b34eeb..08f1e3b8 100644 --- a/src/badger/gui/windows/expandable_message_box.py +++ b/src/badger/gui/windows/expandable_message_box.py @@ -1,17 +1,18 @@ """Error dialog with an expandable "Details" section. Shows a short error message by default; click to reveal the full traceback.""" +from PyQt5.QtCore import Qt +from PyQt5.QtGui import QFont, QFontDatabase, QTextOption from PyQt5.QtWidgets import ( QDialog, - QMessageBox, - QVBoxLayout, QHBoxLayout, QLabel, + QMessageBox, QPushButton, QTextEdit, + QVBoxLayout, ) -from PyQt5.QtGui import QTextOption, QFont, QFontDatabase -from PyQt5.QtCore import Qt + from badger.gui.utils import unset_busy_cursor diff --git a/src/badger/gui/windows/ind_lim_vrange_dialog.py b/src/badger/gui/windows/ind_lim_vrange_dialog.py index d0b591b1..4311b6ab 100644 --- a/src/badger/gui/windows/ind_lim_vrange_dialog.py +++ b/src/badger/gui/windows/ind_lim_vrange_dialog.py @@ -3,25 +3,26 @@ or current range) independently of the global limit settings, providing fine-grained control over the optimization search space.""" -from copy import deepcopy import math +from copy import deepcopy +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( + QComboBox, QDialog, - QWidget, - QHBoxLayout, - QPushButton, - QVBoxLayout, QDoubleSpinBox, + QFrame, QGroupBox, + QHBoxLayout, QLabel, - QComboBox, - QStyledItemDelegate, - QStackedWidget, QLineEdit, - QFrame, + QPushButton, + QStackedWidget, + QStyledItemDelegate, + QVBoxLayout, + QWidget, ) -from PyQt5.QtCore import Qt + from badger.gui.components.bounds_preview import BoundsPreviewBar diff --git a/src/badger/gui/windows/lim_vrange_dialog.py b/src/badger/gui/windows/lim_vrange_dialog.py index 37fa41d1..c717e99b 100644 --- a/src/badger/gui/windows/lim_vrange_dialog.py +++ b/src/badger/gui/windows/lim_vrange_dialog.py @@ -5,23 +5,21 @@ from copy import deepcopy +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( + QComboBox, QDialog, - QWidget, - QHBoxLayout, - QPushButton, - QVBoxLayout, QDoubleSpinBox, - QRadioButton, -) -from PyQt5.QtWidgets import ( QGroupBox, + QHBoxLayout, QLabel, - QComboBox, - QStyledItemDelegate, + QPushButton, + QRadioButton, QStackedWidget, + QStyledItemDelegate, + QVBoxLayout, + QWidget, ) -from PyQt5.QtCore import Qt class BadgerLimitVariableRangeDialog(QDialog): diff --git a/src/badger/gui/windows/load_data_from_run_dialog.py b/src/badger/gui/windows/load_data_from_run_dialog.py index 9a5402bd..e9346c2a 100644 --- a/src/badger/gui/windows/load_data_from_run_dialog.py +++ b/src/badger/gui/windows/load_data_from_run_dialog.py @@ -2,32 +2,34 @@ user browse archived runs, preview their data as a plot, and import selected data points to seed the optimizer with prior observations.""" +from collections.abc import Callable + +import numpy as np +import pandas as pd +import pyqtgraph as pg +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( QDialog, - QWidget, + QFileDialog, QHBoxLayout, - QPushButton, - QVBoxLayout, QLabel, - QFileDialog, QMessageBox, + QPushButton, + QVBoxLayout, + QWidget, ) -from PyQt5.QtCore import Qt -import pyqtgraph as pg -from pyqtgraph.Qt import QtGui, QtCore -from typing import List, Callable -import numpy as np -import pandas as pd +from pyqtgraph.Qt import QtCore, QtGui +from xopt.vocs import VOCS + from badger.archive import ( get_base_run_filename, get_runs, load_run, ) -from badger.gui.components.navigators import HistoryNavigator -from badger.settings import init_settings from badger.errors import BadgerRoutineError +from badger.gui.components.navigators import HistoryNavigator from badger.routine import Routine -from xopt.vocs import VOCS +from badger.settings import init_settings stylesheet_run = """ QPushButton:hover:pressed @@ -56,7 +58,7 @@ def __init__( self, parent: QWidget, env_vocs: VOCS = None, - on_set: Callable[[Routine], None] = None, + on_set: Callable[[Routine], None] | None = None, ): """ Initialize the dialog. @@ -229,7 +231,7 @@ def load_from_file(self) -> None: routine = load_run(file_path) self.preview_run(routine) except Exception as e: - raise BadgerRoutineError(f"{e}") + raise BadgerRoutineError(f"{e}") from e def load_data(self, run_filename: str) -> None: """ @@ -245,7 +247,7 @@ def load_data(self, run_filename: str) -> None: return routine except IndexError: return - except Exception: # failed to load the run + except Exception: # noqa: BLE001 - run load boundary return def init_plots(self) -> None: @@ -298,7 +300,7 @@ def init_plots(self) -> None: return plot_layout def _configure_plot( - self, plot_object: pg.PlotItem, names: List[str] + self, plot_object: pg.PlotItem, names: list[str] ) -> dict[str : pg.PlotCurveItem]: """ Configure the plot with the given data names. @@ -342,7 +344,7 @@ def _configure_plot( return curves def _set_plot_data( - self, names: List[str], curves: dict, data: pd.DataFrame + self, names: list[str], curves: dict, data: pd.DataFrame ) -> None: """ Set data for the plot curves. @@ -376,7 +378,7 @@ def _set_plot_data( hist_x, not_live_data[name].to_numpy(dtype=np.double) ) - def show_vocs_mismatch_dialog(self, list1: List[str], list2: List[str]): + def show_vocs_mismatch_dialog(self, list1: list[str], list2: list[str]): """ Display a helpful dialog notifying the user that the data they are trying to load does not have the same variables and objectives as they have selected in the GUI. diff --git a/src/badger/gui/windows/main_window.py b/src/badger/gui/windows/main_window.py index 6b92fd73..2199edcc 100644 --- a/src/badger/gui/windows/main_window.py +++ b/src/badger/gui/windows/main_window.py @@ -4,9 +4,10 @@ import logging import os from importlib import metadata -from typing import Dict + from PyQt5.QtCore import QThread from PyQt5.QtWidgets import QDesktopWidget, QMainWindow, QMessageBox, QStackedWidget + from badger.gui.components.create_process import CreateProcess from badger.gui.components.process_manager import ProcessManager from badger.gui.pages.home_page import BadgerHomePage @@ -57,7 +58,7 @@ def cleanupThread(self) -> None: if thread in self.thread_list: self.thread_list.remove(thread) - def storeSubprocess(self, process_with_args: Dict) -> None: + def storeSubprocess(self, process_with_args: dict) -> None: logger.info(f"Storing prepared subprocess: {process_with_args}") """ Store the prepared subprocess for later use. @@ -103,7 +104,6 @@ def center(self) -> None: def config_logic(self) -> None: logger.info("Configuring logic.") - pass def closeEvent(self, event) -> None: logger.info("Main window close event triggered.") diff --git a/src/badger/gui/windows/message_dialog.py b/src/badger/gui/windows/message_dialog.py index 788a8911..7a73155c 100644 --- a/src/badger/gui/windows/message_dialog.py +++ b/src/badger/gui/windows/message_dialog.py @@ -2,18 +2,19 @@ with an optional expandable details section in a resizable, scrollable window, used for showing verbose error information or optimization summaries.""" +from PyQt5.QtCore import Qt +from PyQt5.QtGui import QFont, QFontDatabase, QTextOption from PyQt5.QtWidgets import ( QDialog, - QMessageBox, - QVBoxLayout, QHBoxLayout, QLabel, + QMessageBox, QPushButton, QScrollArea, QTextEdit, + QVBoxLayout, ) -from PyQt5.QtGui import QTextOption, QFont, QFontDatabase -from PyQt5.QtCore import Qt + # from ..components.eliding_label import ElidingLabel diff --git a/src/badger/gui/windows/settings_dialog.py b/src/badger/gui/windows/settings_dialog.py index 43a055bb..57c61be0 100644 --- a/src/badger/gui/windows/settings_dialog.py +++ b/src/badger/gui/windows/settings_dialog.py @@ -4,6 +4,7 @@ import logging import os +from typing import ClassVar from PyQt5.QtCore import Qt @@ -29,8 +30,8 @@ class BadgerSettingsDialog(QDialog): - theme_list = ["default", "light", "dark"] - theme_idx_dict = { + theme_list: ClassVar[list[str]] = ["default", "light", "dark"] + theme_idx_dict: ClassVar[dict[str, int]] = { "default": 0, "light": 1, "dark": 2, @@ -277,7 +278,7 @@ def restore_settings(self): if theme_prev != theme_curr: self.set_theme(theme_prev) - for key in self.settings.keys(): + for key in self.settings: self.config_singleton.write_value(key, self.settings[key]["value"]) self.reject() diff --git a/src/badger/gui/windows/terminition_condition_dialog.py b/src/badger/gui/windows/terminition_condition_dialog.py index 57a1c0a9..d0e8e640 100644 --- a/src/badger/gui/windows/terminition_condition_dialog.py +++ b/src/badger/gui/windows/terminition_condition_dialog.py @@ -3,23 +3,20 @@ number of evaluations or after a maximum elapsed time.""" from PyQt5.QtWidgets import ( + QComboBox, QDialog, - QWidget, - QHBoxLayout, - QPushButton, - QVBoxLayout, - QSpinBox, QDoubleSpinBox, -) -from PyQt5.QtWidgets import ( QGroupBox, + QHBoxLayout, QLabel, - QComboBox, - QStyledItemDelegate, + QPushButton, + QSpinBox, QStackedWidget, + QStyledItemDelegate, + QVBoxLayout, + QWidget, ) - stylesheet_run = """ QPushButton:hover:pressed { diff --git a/src/badger/gui/windows/var_dialog.py b/src/badger/gui/windows/var_dialog.py index 247e96b1..a69f3cf5 100644 --- a/src/badger/gui/windows/var_dialog.py +++ b/src/badger/gui/windows/var_dialog.py @@ -2,17 +2,23 @@ for finding and adding environment variables by name, querying the environment plugin to verify the variable exists and retrieve its current value.""" +import logging + from PyQt5.QtWidgets import ( QDialog, - QWidget, + QGroupBox, QHBoxLayout, QLineEdit, + QMessageBox, QPushButton, QVBoxLayout, + QWidget, ) -from PyQt5.QtWidgets import QGroupBox, QMessageBox -from badger.gui.components.labeled_lineedit import labeled_lineedit + from badger.environment import instantiate_env +from badger.gui.components.labeled_lineedit import labeled_lineedit + +logger = logging.getLogger(__name__) class BadgerVariableDialog(QDialog): @@ -98,6 +104,7 @@ def check_var(self): self.btn_add.setDisabled(False) except Exception: + logger.exception(f"Variable {name} lookup failed") self.edit_value.edit.setText("") self.edit_min.edit.setText("") self.edit_max.edit.setText("") diff --git a/src/badger/interface.py b/src/badger/interface.py index 63ad12ab..902965fd 100644 --- a/src/badger/interface.py +++ b/src/badger/interface.py @@ -4,7 +4,7 @@ import pickle from abc import ABC, abstractmethod -from typing import Any, ClassVar, Dict, List, TypedDict +from typing import Any, ClassVar, TypedDict from pydantic import BaseModel @@ -14,7 +14,7 @@ def log(func): def func_log(*args, **kwargs): if func.__name__ == "set_values": - if "channel_inputs" in kwargs.keys(): + if "channel_inputs" in kwargs: channel_inputs = kwargs["channel_inputs"] else: channel_inputs = args[1] @@ -45,8 +45,8 @@ def func_log(*args, **kwargs): class InterfaceInfo(TypedDict): - interface: Dict[str, str] - vars: List[Dict[str, str]] + interface: dict[str, str] + vars: list[dict[str, str]] class Interface(BaseModel, ABC): @@ -55,7 +55,7 @@ class Interface(BaseModel, ABC): # params: float = Field(..., description='Example intf parameter') # Private variables - _logs: List[Dict] = [] # TODO: Add a property for it? + _logs: list[dict] = [] # TODO: Add a property for it? def start_recording(self): self._logs = [] @@ -82,12 +82,12 @@ def dump_recording(self, filename): # Environment should only call this method to get channels @abstractmethod - def get_values(self, channel_names: List[str]) -> Dict[str, Any]: + def get_values(self, channel_names: list[str]) -> dict[str, Any]: pass # Environment should only call this method to set channels @abstractmethod - def set_values(self, channel_inputs: Dict[str, Any]): + def set_values(self, channel_inputs: dict[str, Any]): pass def reset_interface(self): @@ -95,7 +95,6 @@ def reset_interface(self): Called after the application forks (i.e. after spawning a new multiprocess.Process) Subclasses should use this to reset any undesirable process wide state """ - pass def get_value(self, channel_name: str, **kwargs) -> Any: return self.get_values([channel_name], **kwargs)[channel_name] @@ -103,7 +102,7 @@ def get_value(self, channel_name: str, **kwargs) -> Any: def set_value(self, channel_name: str, channel_value, **kwargs): return self.set_values({channel_name: channel_value}, **kwargs) - def get_info(self, channels: List[str]) -> InterfaceInfo: + def get_info(self, channels: list[str]) -> InterfaceInfo: """ Optional; Returns information about the channels and environment for display diff --git a/src/badger/log.py b/src/badger/log.py index bf0b961d..6f76062f 100644 --- a/src/badger/log.py +++ b/src/badger/log.py @@ -9,15 +9,16 @@ For example usage (in a simple context), see src/badger/tests/test_multiprocess_logging.py """ -import os -import datetime -import logging -import atexit +from __future__ import annotations +import atexit +import logging +import os +from datetime import UTC, datetime from logging.handlers import QueueHandler, QueueListener from multiprocessing import Queue -from badger.settings import get_user_config_folder -from badger.settings import init_settings + +from badger.settings import get_user_config_folder, init_settings logger = logging.getLogger(__name__) @@ -157,7 +158,7 @@ def get_logfile_name(self): Get name of the logfile for today, which is in form of: "log__.log" """ - today = datetime.date.today() + today = datetime.now(tz=UTC).date() log_filename = f"log_{today.year:04d}_{today.month:02d}_{today.day:02d}.log" return log_filename @@ -195,11 +196,11 @@ def create_log_dir(self, log_dir: str): def configure_process_logging( - log_queue: Queue = None, + log_queue: Queue | None = None, logger_name: str = "badger", log_level: str = "DEBUG", - process_name: str = None, -): + process_name: str | None = None, +) -> None: """ Configure logging in a process to send logs to the shared queue. This should be ran in all subprocesses b4 logging, and also in the main process where diff --git a/src/badger/logbook.py b/src/badger/logbook.py index b46477d6..d57a4a94 100644 --- a/src/badger/logbook.py +++ b/src/badger/logbook.py @@ -2,13 +2,13 @@ XML entry with run summary (gain, duration, algorithm) and attaches a screenshot of the GUI for the facility archive.""" -import os -from datetime import datetime import logging +import os +from datetime import UTC, datetime -from badger.settings import init_settings from badger.archive import BADGER_ARCHIVE_ROOT from badger.errors import BadgerConfigError, BadgerLogbookError +from badger.settings import init_settings logger = logging.getLogger(__name__) @@ -23,8 +23,8 @@ def send_to_logbook(routine, widget=None): - from xml.etree import ElementTree from re import sub + from xml.etree import ElementTree log_text = "" routine_name = routine.name @@ -50,13 +50,11 @@ def send_to_logbook(routine, widget=None): log_text += f"Environment name: {env_name}\n" log_text += f"Optimization algorithm: {generator_name}\n" log_text += f"Data location: {data_path}\n" - try: - log_text += f"Log location: {BADGER_LOGBOOK_ROOT}\n" - except: - pass + + log_text += f"Log location: {BADGER_LOGBOOK_ROOT}\n" # Generate the xml data - curr_time = datetime.now() + curr_time = datetime.now(tz=UTC) if os.name == "nt": timestr = curr_time.strftime("%Y-%m-%dT%H%M%S") else: @@ -92,13 +90,12 @@ def send_to_logbook(routine, widget=None): fileName = os.path.join(BADGER_LOGBOOK_ROOT, metainfo.text) fileName = fileName.rstrip(".xml") - xmlFile = open(fileName + ".xml", "w") - rawString = ElementTree.tostring(log_entry, "utf-8").decode("utf-8") - parsedString = sub(r"(?=<[^/].*>)", "\n", rawString) - xmlString = parsedString[1:] - xmlFile.write(xmlString) - xmlFile.write("\n") # Close with newline so cron job parses correctly - xmlFile.close() + with open(fileName + ".xml", "w") as xmlFile: + rawString = ElementTree.tostring(log_entry, "utf-8").decode("utf-8") + parsedString = sub(r"(?=<[^/].*>)", "\n", rawString) + xmlString = parsedString[1:] + xmlFile.write(xmlString) + xmlFile.write("\n") # Close with newline so cron job parses correctly screenshot(widget, f"{fileName}.png") diff --git a/src/badger/logger/__init__.py b/src/badger/logger/__init__.py index db46bd78..659544cc 100644 --- a/src/badger/logger/__init__.py +++ b/src/badger/logger/__init__.py @@ -2,12 +2,11 @@ and to a JSON file (JSONLogger). Both subscribe to lifecycle events (start, step, end) from the optimization loop.""" -from __future__ import print_function -import os import json +import os -from badger.logger.observer import _Tracker from badger.logger.event import Events +from badger.logger.observer import _Tracker from badger.logger.util import Colours @@ -22,7 +21,7 @@ class ScreenLogger(_Tracker): def __init__(self, verbose=2): self._verbose = verbose self._header_length = None - super(ScreenLogger, self).__init__() + super().__init__() @property def verbose(self): @@ -128,7 +127,7 @@ def __init__(self, path, reset=True): os.remove(self._path) except OSError: pass - super(JSONLogger, self).__init__() + super().__init__() def update(self, event, solution): if event == Events.OPTIMIZATION_STEP: diff --git a/src/badger/logger/observer.py b/src/badger/logger/observer.py index 8526297b..72c39c50 100644 --- a/src/badger/logger/observer.py +++ b/src/badger/logger/observer.py @@ -1,7 +1,8 @@ """Base classes for logger observers. Observer defines the update() interface; _Tracker adds iteration counting and elapsed-time tracking.""" -from datetime import datetime +from datetime import UTC, datetime + from badger.logger.event import Events @@ -10,7 +11,7 @@ def update(self, event, solution): raise NotImplementedError -class _Tracker(object): +class _Tracker: def __init__(self): self._iterations = 0 @@ -22,7 +23,7 @@ def _update_tracker(self, event, solution): self._iterations += 1 def _time_metrics(self): - now = datetime.now() + now = datetime.now(tz=UTC) if self._start_time is None: self._start_time = now if self._previous_time is None: diff --git a/src/badger/routine.py b/src/badger/routine.py index cca84e8e..843eb1c7 100644 --- a/src/badger/routine.py +++ b/src/badger/routine.py @@ -11,59 +11,61 @@ state and generate initial sampling points. """ -import logging import json +import logging from copy import deepcopy -from typing import Any, List, Optional +from typing import Any + import numpy as np import pandas as pd + +# Import xopt.generators at startup so they don't need to be imported +# each time a Routine is created +import xopt.generators.bayesian +import xopt.generators.sequential # noqa: F401 from pandas import DataFrame from pydantic import ( ConfigDict, Field, - field_validator, - model_validator, SerializeAsAny, ValidationInfo, + field_validator, + model_validator, ) -from xopt import Evaluator, VOCS, Xopt +from xopt import VOCS, Evaluator, Xopt from xopt.generators import get_generator -from xopt.vocs import get_local_region from xopt.generators.sequential import SequentialGenerator -from badger.utils import curr_ts +from xopt.vocs import get_local_region + from badger.environment import BaseEnvironment, instantiate_env from badger.factory import get_env - -# Import xopt.generators at startup so they don't need to be imported -# each time a Routine is created -import xopt.generators.bayesian # noqa: F401 -import xopt.generators.sequential # noqa: F401 +from badger.utils import curr_ts logger = logging.getLogger(__name__) class Routine(Xopt): - id: Optional[str] = Field(None) - creation_ts: Optional[str] = Field(None) # Timestamp of routine creation + id: str | None = Field(None) + creation_ts: str | None = Field(None) # Timestamp of routine creation name: str - description: Optional[str] = Field(None) + description: str | None = Field(None) environment: SerializeAsAny[BaseEnvironment] - initial_points: Optional[DataFrame] = Field(None) - critical_constraint_names: Optional[List[str]] = Field([]) - tags: Optional[List] = Field(None) - script: Optional[str] = Field(None) + initial_points: DataFrame | None = Field(None) + critical_constraint_names: list[str] | None = Field([]) + tags: list | None = Field(None) + script: str | None = Field(None) # Store relative to current params - relative_to_current: Optional[bool] = Field(False) - vrange_limit_options: Optional[dict] = Field(None) - vrange_hard_limit: Optional[dict] = Field({}) # override hard limits - initial_point_actions: Optional[List] = Field(None) - additional_variables: Optional[List[str]] = Field([]) - formulas: Optional[dict[str, dict[str, Any]]] = Field({}) - constraint_formulas: Optional[dict[str, dict[str, Any]]] = Field({}) - observable_formulas: Optional[dict[str, dict[str, Any]]] = Field({}) + relative_to_current: bool | None = Field(False) + vrange_limit_options: dict | None = Field(None) + vrange_hard_limit: dict | None = Field({}) # override hard limits + initial_point_actions: list | None = Field(None) + additional_variables: list[str] | None = Field([]) + formulas: dict[str, dict[str, Any]] | None = Field({}) + constraint_formulas: dict[str, dict[str, Any]] | None = Field({}) + observable_formulas: dict[str, dict[str, Any]] | None = Field({}) # Other meta data - badger_version: Optional[str] = Field(None) - xopt_version: Optional[str] = Field(None) + badger_version: str | None = Field(None) + xopt_version: str | None = Field(None) model_config = ConfigDict(arbitrary_types_allowed=True) @@ -108,24 +110,23 @@ def validate_model(cls, data: Any): data["generator"].is_active = False # validate data (if it exists - if "data" in data: - if isinstance(data["data"], dict): - logger.debug("Validating and converting data to DataFrame.") - try: - data["data"] = pd.DataFrame(data["data"]) - except IndexError: - data["data"] = pd.DataFrame(data["data"], index=[0]) - - data["data"].index = data["data"].index.astype(int) - data["data"].sort_index(inplace=True) - - # Add data one row at a time to avoid generator issues - if isinstance(data["generator"], SequentialGenerator): - logger.debug("Setting data for SequentialGenerator.") - data["generator"].set_data(data["data"]) - else: - logger.debug("Adding data to generator.") - data["generator"].add_data(data["data"]) + if "data" in data and isinstance(data["data"], dict): + logger.debug("Validating and converting data to DataFrame.") + try: + data["data"] = pd.DataFrame(data["data"]) + except IndexError: + data["data"] = pd.DataFrame(data["data"], index=[0]) + + data["data"].index = data["data"].index.astype(int) + data["data"].sort_index(inplace=True) + + # Add data one row at a time to avoid generator issues + if isinstance(data["generator"], SequentialGenerator): + logger.debug("Setting data for SequentialGenerator.") + data["generator"].set_data(data["data"]) + else: + logger.debug("Adding data to generator.") + data["generator"].add_data(data["data"]) # instantiate env if isinstance(data["environment"], dict): diff --git a/src/badger/settings.py b/src/badger/settings.py index 29639390..7eab7aa5 100644 --- a/src/badger/settings.py +++ b/src/badger/settings.py @@ -13,7 +13,7 @@ import platform import shutil from importlib import resources -from typing import Any, Dict, Optional, Union +from typing import Any import yaml from pydantic import BaseModel, Field, ValidationError @@ -41,7 +41,7 @@ class Setting(BaseModel): display_name: str description: str - value: Optional[Union[str, int, bool, None]] = Field( + value: None | str | int | bool = Field( None, description="The value of the setting which can be of different types." ) is_path: bool @@ -154,9 +154,9 @@ class BadgerConfig(BaseModel): class ConfigSingleton: _instance = None - def __new__(cls, config_path: str = None, user_flag: bool = False): + def __new__(cls, config_path: str | None = None, user_flag: bool = False): if cls._instance is None: - cls._instance = super(ConfigSingleton, cls).__new__(cls) + cls._instance = super().__new__(cls) cls._instance.user_flag = user_flag cls._instance._config = cls.load_or_create_config(config_path) cls._instance.config_path = config_path @@ -221,7 +221,7 @@ def load_or_create_config(cls, config_path: str) -> BadgerConfig: def config(self) -> BadgerConfig: return self._config - def update_and_save_config(self, updates: Dict[str, Any]) -> None: + def update_and_save_config(self, updates: dict[str, Any]) -> None: """Saves changes to the config file. Parameters @@ -242,7 +242,7 @@ def update_and_save_config(self, updates: Dict[str, Any]) -> None: self._config = BadgerConfig(**config_data) def _update_config_by_dot_key( - self, config_data: Dict[str, Any], dot_key: str, value: Any + self, config_data: dict[str, Any], dot_key: str, value: Any ) -> None: """Update the config data with the provided value using dot-separated keys.""" keys = dot_key.split(":") @@ -255,7 +255,7 @@ def _update_config_by_dot_key( else: d[last_key] = value - def list_settings(self) -> Dict[str, Any]: + def list_settings(self) -> dict[str, Any]: """List all the settings in Badger Returns @@ -266,7 +266,7 @@ def list_settings(self) -> Dict[str, Any]: """ return self._config.model_dump(by_alias=True) - def list_path_settings(self) -> Dict[str, Any]: + def list_path_settings(self) -> dict[str, Any]: """List all the path-related settings in Badger Returns @@ -411,8 +411,8 @@ def write_value(self, key: str, value: Any) -> None: updates = {} sub_dict = updates - for key in keys[:-1]: - sub_dict = sub_dict.setdefault(key, {}) + for k in keys[:-1]: + sub_dict = sub_dict.setdefault(k, {}) sub_dict[keys[-1]] = value logger.info(f"writing to config file, setting: {key} = {value}") @@ -427,13 +427,13 @@ def reset_settings(self) -> None: ) -def init_settings(config_arg: str = None) -> ConfigSingleton: +def init_settings(config_arg: str | None = None) -> ConfigSingleton: """ Builds and returns an instance of the ConfigSingleton class. Parameters ---------- - config_arg: str + config_arg: str | None a path to a config file passed through the --config__filepath argument Returns @@ -575,7 +575,7 @@ def mock_settings(): config_singleton.write_value("BADGER_TEMP_DIRECTORY", temp_dir) # Set other settings to the default values - for key in config_singleton.config.model_dump(by_alias=True).keys(): + for key in config_singleton.config.model_dump(by_alias=True): config_singleton.write_value( key, config_singleton.config.model_dump(by_alias=True)[key]["value"] ) diff --git a/src/badger/tests/mock/plugins/environments/multiobjective_test/__init__.py b/src/badger/tests/mock/plugins/environments/multiobjective_test/__init__.py index 5029ccc6..2b074eb6 100644 --- a/src/badger/tests/mock/plugins/environments/multiobjective_test/__init__.py +++ b/src/badger/tests/mock/plugins/environments/multiobjective_test/__init__.py @@ -1,12 +1,15 @@ +from typing import ClassVar + import torch + from badger import environment from badger.errors import BadgerNoInterfaceError class Environment(environment.Environment): name = "multiobjective_test" - variables = {f"x{i}": [-1, 1] for i in range(20)} - observables = ["f1", "f2"] + variables: ClassVar[dict[str, list[float]]] = {f"x{i}": [-1, 1] for i in range(20)} + observables: ClassVar[list[str]] = ["f1", "f2"] flag: int = 0 diff --git a/src/badger/tests/mock/plugins/environments/test/__init__.py b/src/badger/tests/mock/plugins/environments/test/__init__.py index a871c4ec..9d76c265 100644 --- a/src/badger/tests/mock/plugins/environments/test/__init__.py +++ b/src/badger/tests/mock/plugins/environments/test/__init__.py @@ -1,13 +1,16 @@ +import time +from typing import ClassVar + import torch + from badger import environment from badger.errors import BadgerNoInterfaceError -import time class Environment(environment.Environment): name = "test" - variables = {f"x{i}": [-1, 1] for i in range(20)} - observables = ["f", "c"] + variables: ClassVar[dict[str, list[float]]] = {f"x{i}": [-1, 1] for i in range(20)} + observables: ClassVar[list[str]] = ["f", "c"] flag: int = 0 delay: float = 0.0 diff --git a/src/badger/tests/mock/plugins/interfaces/test/__init__.py b/src/badger/tests/mock/plugins/interfaces/test/__init__.py index e60cbc6b..8bae3784 100644 --- a/src/badger/tests/mock/plugins/interfaces/test/__init__.py +++ b/src/badger/tests/mock/plugins/interfaces/test/__init__.py @@ -1,5 +1,4 @@ from badger import interface -from typing import Dict class Interface(interface.Interface): @@ -7,7 +6,7 @@ class Interface(interface.Interface): flag: int = 0 # Private variables - _states: Dict + _states: dict def __init__(self, **data): super().__init__(**data) diff --git a/src/badger/tests/multiprocess_logging.py b/src/badger/tests/multiprocess_logging.py index 1eccb0d6..a6afea19 100644 --- a/src/badger/tests/multiprocess_logging.py +++ b/src/badger/tests/multiprocess_logging.py @@ -1,7 +1,8 @@ import logging import time from multiprocessing import Process -from badger.log import get_logging_manager, configure_process_logging + +from badger.log import configure_process_logging, get_logging_manager """ Note: this test does not have "test" in the filename so pytest does not discover it when running tests in a group. diff --git a/src/badger/tests/test_cli_basic.py b/src/badger/tests/test_cli_basic.py index 1ceaf1bc..c61d20e1 100644 --- a/src/badger/tests/test_cli_basic.py +++ b/src/badger/tests/test_cli_basic.py @@ -12,7 +12,7 @@ def capture(command): def test_cli_main(): command = ["badger"] - out, err, exitcode = capture(command) + out, _, exitcode = capture(command) assert exitcode == 0 @@ -36,7 +36,7 @@ def test_list_algo(): from badger.factory import ALGO_EXCLUDED command = ["badger", "generator"] - out, err, exitcode = capture(command) + out, _, exitcode = capture(command) assert exitcode == 0 diff --git a/src/badger/tests/test_core.py b/src/badger/tests/test_core.py index 3925b516..bc2127e4 100644 --- a/src/badger/tests/test_core.py +++ b/src/badger/tests/test_core.py @@ -1,6 +1,8 @@ import os -import pytest + import pandas as pd +import pytest + from badger.errors import BadgerRunTerminated @@ -11,7 +13,7 @@ class TestCore: @pytest.fixture(autouse=True, scope="function") def test_core_setup(self, *args, **kwargs) -> None: - super(TestCore, self).__init__(*args, **kwargs) + super().__init__(*args, **kwargs) self.count = 0 self.candidates = None self.points_eval_list = [] diff --git a/src/badger/tests/test_core_subprocess.py b/src/badger/tests/test_core_subprocess.py index 972a3c0b..5e20cbf6 100644 --- a/src/badger/tests/test_core_subprocess.py +++ b/src/badger/tests/test_core_subprocess.py @@ -36,7 +36,7 @@ def init_multiprocessing(self): @pytest.fixture(autouse=True, scope="function") def test_core_setup(self, *args, **kwargs) -> None: - super(TestCore, self).__init__(*args, **kwargs) + super().__init__(*args, **kwargs) self.count = 0 self.candidates = None self.points_eval_list = [] diff --git a/src/badger/tests/test_env.py b/src/badger/tests/test_env.py index 37e4b974..6f783f4e 100644 --- a/src/badger/tests/test_env.py +++ b/src/badger/tests/test_env.py @@ -2,7 +2,7 @@ def test_find_env(): - from badger.factory import list_env, get_env + from badger.factory import get_env, list_env assert len(list_env()) == 2 @@ -61,10 +61,10 @@ def test_list_observables(): def test_get_variables(): - from badger.factory import get_env, get_intf from badger.errors import ( BadgerNoInterfaceError, ) + from badger.factory import get_env, get_intf Interface, _ = get_intf("test") intf = Interface() @@ -93,11 +93,11 @@ def test_get_variables(): def testset_variables(): - from badger.factory import get_env, get_intf from badger.errors import ( BadgerEnvVarError, BadgerNoInterfaceError, ) + from badger.factory import get_env, get_intf Interface, _ = get_intf("test") intf = Interface() @@ -141,8 +141,8 @@ def testset_variables(): def testget_observables(): - from badger.factory import get_env, get_intf from badger.errors import BadgerNoInterfaceError + from badger.factory import get_env, get_intf Interface, _ = get_intf("test") intf = Interface() diff --git a/src/badger/tests/test_environment.py b/src/badger/tests/test_environment.py index ca14f1b9..ffeeabc4 100644 --- a/src/badger/tests/test_environment.py +++ b/src/badger/tests/test_environment.py @@ -1,9 +1,10 @@ import json -import pytest -import numpy as np -from typing import Dict, List +from typing import ClassVar from unittest.mock import Mock +import numpy as np +import pytest + from badger.environment import BaseEnvironment, Environment from badger.errors import BadgerEnvVarError, BadgerNoInterfaceError from badger.interface import Interface @@ -28,18 +29,20 @@ def test_basic_env(self): class TestEnv(BaseEnvironment): name = "test" - variables = {f"x{i}": [-1, 1] for i in range(20)} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = { + f"x{i}": [-1, 1] for i in range(20) + } + observables: ClassVar[list[str]] = ["f"] my_flag: int = 0 - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.5 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: return {ele: 1.0 for ele in observable_names} env = TestEnv() @@ -63,8 +66,8 @@ def test_environment_with_interface(self): class TestEnv(Environment): name = "test" - variables = {"x1": [-1, 1], "x2": [-2, 2]} - observables = ["f", "g"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1], "x2": [-2, 2]} + observables: ClassVar[list[str]] = ["f", "g"] env = TestEnv(interface=mock_interface) @@ -90,8 +93,8 @@ def test_environment_no_interface_error(self): class TestEnv(Environment): name = "test" - variables = {"x1": [-1, 1]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1]} + observables: ClassVar[list[str]] = ["f"] env = TestEnv() @@ -109,19 +112,19 @@ def test_bounds_validation(self): class TestEnv(BaseEnvironment): name = "test" - variables = {"x1": [-1, 1], "x2": [0, 10]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1], "x2": [0, 10]} + observables: ClassVar[list[str]] = ["f"] - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.0 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: return {name: 1.0 for name in observable_names} - def get_bounds(self, variable_names: List[str]) -> dict[str, float]: + def get_bounds(self, variable_names: list[str]) -> dict[str, float]: # Test invalid bounds - should be caught by decorator return {"x1": [1, -1]} # upper < lower @@ -137,16 +140,16 @@ def test_setpoint_validation(self): class TestEnv(BaseEnvironment): name = "test" - variables = {"x1": [-1, 1], "x2": [0, 10]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1], "x2": [0, 10]} + observables: ClassVar[list[str]] = ["f"] - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.0 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: return {name: 1.0 for name in observable_names} env = TestEnv() @@ -170,16 +173,16 @@ def test_formula_processing(self): class TestEnv(BaseEnvironment): name = "test" - variables = {"x1": [-1, 1], "x2": [-1, 1]} - observables = ["f", "g", "h"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1], "x2": [-1, 1]} + observables: ClassVar[list[str]] = ["f", "g", "h"] - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.0 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: # Return base observables return { name: {"f": 2.0, "g": 3.0, "h": 1.0}[name] @@ -204,16 +207,16 @@ def test_convenience_methods(self): class TestEnv(BaseEnvironment): name = "test" - variables = {"x1": [-1, 1]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1]} + observables: ClassVar[list[str]] = ["f"] - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.5 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): self._last_set = variable_inputs - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: return {name: 1.5 for name in observable_names} env = TestEnv() @@ -234,8 +237,12 @@ def test_variable_names_property(self): class TestEnv(Environment): name = "test" - variables = {"x1": [-1, 1], "x2": [0, 10], "x3": [-5, 5]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = { + "x1": [-1, 1], + "x2": [0, 10], + "x3": [-5, 5], + } + observables: ClassVar[list[str]] = ["f"] env = TestEnv(interface=mock_interface) assert set(env.variable_names) == {"x1", "x2", "x3"} @@ -245,16 +252,16 @@ def test_get_system_states(self): class TestEnv(BaseEnvironment): name = "test" - variables = {"x1": [-1, 1]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1]} + observables: ClassVar[list[str]] = ["f"] - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.0 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: return {name: 1.0 for name in observable_names} env = TestEnv() @@ -265,16 +272,16 @@ def test_search_not_implemented(self): class TestEnv(BaseEnvironment): name = "test" - variables = {"x1": [-1, 1]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1]} + observables: ClassVar[list[str]] = ["f"] - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.0 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: return {name: 1.0 for name in observable_names} env = TestEnv() @@ -290,18 +297,20 @@ def test_env_in_routine(self): class TestEnv(BaseEnvironment): name = "test" - variables = {f"x{i}": [-1, 1] for i in range(20)} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = { + f"x{i}": [-1, 1] for i in range(20) + } + observables: ClassVar[list[str]] = ["f"] my_flag: int = 0 - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.0 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: return {ele: 1.0 for ele in observable_names} env = TestEnv() @@ -344,16 +353,16 @@ def test_complex_formula_processing(self): class TestEnv(BaseEnvironment): name = "test" - variables = {"x1": [-1, 1], "x2": [-1, 1]} - observables = ["data", "noise", "signal"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1], "x2": [-1, 1]} + observables: ClassVar[list[str]] = ["data", "noise", "signal"] - def get_variables(self, variable_names: List[str]) -> Dict[str, float]: + def get_variables(self, variable_names: list[str]) -> dict[str, float]: return {name: 0.0 for name in variable_names} - def set_variables(self, variable_inputs: Dict[str, float]): + def set_variables(self, variable_inputs: dict[str, float]): pass - def get_observables(self, observable_names: List[str]) -> Dict[str, float]: + def get_observables(self, observable_names: list[str]) -> dict[str, float]: # Return mock data for different observables data_map = { "data": np.array([1, 2, 3, 4, 5]), @@ -390,8 +399,8 @@ def test_instantiate_env_function(self): class TestEnv(Environment): name = "test" - variables = {"x1": [-1, 1]} - observables = ["f"] + variables: ClassVar[dict[str, list[float]]] = {"x1": [-1, 1]} + observables: ClassVar[list[str]] = ["f"] test_param: float = 1.0 # Test with interface config diff --git a/src/badger/tests/test_factory.py b/src/badger/tests/test_factory.py index ab4d3e3a..c9ebf368 100644 --- a/src/badger/tests/test_factory.py +++ b/src/badger/tests/test_factory.py @@ -39,7 +39,7 @@ def test_generator_generation(self): elif name == "latin_hypercube": test_vocs = deepcopy(TEST_VOCS_BASE) test_vocs.objectives = { - k: ExploreObjective() for k in test_vocs.objectives.keys() + k: ExploreObjective() for k in test_vocs.objectives } gen_class(vocs=test_vocs, **gen_config) elif name == "bax": diff --git a/src/badger/tests/test_formulas.py b/src/badger/tests/test_formulas.py index 3d9fd088..ff1da129 100644 --- a/src/badger/tests/test_formulas.py +++ b/src/badger/tests/test_formulas.py @@ -1,11 +1,12 @@ -import pytest import numpy as np +import pytest + from badger.formula import ( - safe_var_name, extract_variable_keys, find_used_names, - suggest_name, interpret_expression, + safe_var_name, + suggest_name, ) diff --git a/src/badger/tests/test_gui_basic.py b/src/badger/tests/test_gui_basic.py index d094ab1d..bfd3711a 100644 --- a/src/badger/tests/test_gui_basic.py +++ b/src/badger/tests/test_gui_basic.py @@ -74,7 +74,7 @@ def test_close_main(qtbot, init_multiprocessing): window.close() # this action should release the env # So we expect an AttributeError here with pytest.raises(AttributeError): - home_page.run_monitor.routine.environment + _ = home_page.run_monitor.routine.environment window.process_manager.close_proccesses() @@ -237,10 +237,11 @@ def test_default_low_noise_prior_in_bo(qtbot, init_multiprocessing): def test_default_turbo_in_bo(qtbot): return + import yaml + from xopt.generators import all_generator_names + from badger.gui.windows.main_window import BadgerMainWindow from badger.tests.utils import fix_db_path_issue - from xopt.generators import all_generator_names - import yaml fix_db_path_issue() diff --git a/src/badger/tests/test_lib_basic.py b/src/badger/tests/test_lib_basic.py index df179639..5e59aecb 100644 --- a/src/badger/tests/test_lib_basic.py +++ b/src/badger/tests/test_lib_basic.py @@ -6,7 +6,7 @@ def test_env_api(): - from badger.factory import list_env, get_env + from badger.factory import get_env, list_env assert len(list_env()) == 2 @@ -15,7 +15,7 @@ def test_env_api(): def test_intf_api(): - from badger.factory import list_intf, get_intf + from badger.factory import get_intf, list_intf assert len(list_intf()) == 1 diff --git a/src/badger/tests/test_routine_runner.py b/src/badger/tests/test_routine_runner.py index 1f9648a5..f4b3fe0b 100644 --- a/src/badger/tests/test_routine_runner.py +++ b/src/badger/tests/test_routine_runner.py @@ -1,10 +1,10 @@ import multiprocessing +from unittest.mock import Mock import pytest from PyQt5.QtCore import QEventLoop, Qt, QTimer from PyQt5.QtTest import QSignalSpy from PyQt5.QtWidgets import QApplication -from unittest.mock import Mock class TestRoutineRunner: diff --git a/src/badger/tests/test_settings.py b/src/badger/tests/test_settings.py index 55978f06..04c2c66e 100644 --- a/src/badger/tests/test_settings.py +++ b/src/badger/tests/test_settings.py @@ -1,14 +1,16 @@ import os import platform +from unittest.mock import MagicMock, mock_open, patch + import pytest import yaml -from unittest.mock import patch, MagicMock, mock_open + from badger.settings import ( - init_settings, - get_user_config_folder, - ConfigSingleton, BadgerConfig, + ConfigSingleton, Setting, + get_user_config_folder, + init_settings, ) @@ -101,170 +103,186 @@ def test_get_user_config_folder_unsupported_os(self, monkeypatch): def test_init_settings(self): mock_config_singleton = MagicMock(spec=ConfigSingleton) - with patch( - "badger.settings.get_user_config_folder", return_value="/mock/config/folder" + with ( + patch( + "badger.settings.get_user_config_folder", + return_value="/mock/config/folder", + ), + patch( + "badger.settings.ConfigSingleton", return_value=mock_config_singleton + ) as mock_config_cls, ): with patch( - "badger.settings.ConfigSingleton", return_value=mock_config_singleton - ) as mock_config_cls: - with patch( - "badger.settings.get_or_create_temp_directory" - ) as mock_get_or_create_temp_directory: - config_singleton = init_settings() - mock_config_cls.assert_called_once_with( - "/mock/config/folder/config.yaml", False - ) - mock_get_or_create_temp_directory.assert_called_once_with( - mock_config_singleton - ) - assert config_singleton == mock_config_singleton + "badger.settings.get_or_create_temp_directory" + ) as mock_get_or_create_temp_directory: + config_singleton = init_settings() + mock_config_cls.assert_called_once_with( + "/mock/config/folder/config.yaml", False + ) + mock_get_or_create_temp_directory.assert_called_once_with( + mock_config_singleton + ) + assert config_singleton == mock_config_singleton def test_config_singleton_initialization(self): # Patch the open function and ensure it reads the correct mock configuration - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - assert ( - config_singleton.config.BADGER_PLUGIN_ROOT.value - == "/mock/plugin/root" - ) - assert ( - config_singleton.config.BADGER_TEMPLATE_ROOT.value - == "/mock/template/root" - ) + ), + ): + config_singleton = ConfigSingleton(self.config_file) + assert ( + config_singleton.config.BADGER_PLUGIN_ROOT.value == "/mock/plugin/root" + ) + assert ( + config_singleton.config.BADGER_TEMPLATE_ROOT.value + == "/mock/template/root" + ) def test_config_singleton_create_new_config_if_not_exists(self): - with patch("os.path.exists", return_value=False): - with patch( - "badger.settings.BadgerConfig", return_value=self.mock_badger_config - ): - config_singleton = ConfigSingleton(self.config_file) - assert isinstance(config_singleton.config, BadgerConfig) + with ( + patch("os.path.exists", return_value=False), + patch("badger.settings.BadgerConfig", return_value=self.mock_badger_config), + ): + config_singleton = ConfigSingleton(self.config_file) + assert isinstance(config_singleton.config, BadgerConfig) def test_update_and_save_config(self): - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - updates = {"BADGER_PLUGIN_ROOT": "/new/plugin/root"} - config_singleton.update_and_save_config(updates) - assert ( - config_singleton.config.BADGER_PLUGIN_ROOT.value - == "/new/plugin/root" - ) + ), + ): + config_singleton = ConfigSingleton(self.config_file) + updates = {"BADGER_PLUGIN_ROOT": "/new/plugin/root"} + config_singleton.update_and_save_config(updates) + assert ( + config_singleton.config.BADGER_PLUGIN_ROOT.value == "/new/plugin/root" + ) def test_list_settings(self): - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - settings = config_singleton.list_settings() - assert "BADGER_PLUGIN_ROOT" in settings - assert settings["BADGER_PLUGIN_ROOT"]["value"] == "/mock/plugin/root" + ), + ): + config_singleton = ConfigSingleton(self.config_file) + settings = config_singleton.list_settings() + assert "BADGER_PLUGIN_ROOT" in settings + assert settings["BADGER_PLUGIN_ROOT"]["value"] == "/mock/plugin/root" def test_read_value(self): - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - value = config_singleton.read_value("BADGER_PLUGIN_ROOT") - assert value == "/mock/plugin/root" + ), + ): + config_singleton = ConfigSingleton(self.config_file) + value = config_singleton.read_value("BADGER_PLUGIN_ROOT") + assert value == "/mock/plugin/root" def test_read_display_name(self): - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - display_name = config_singleton.read_display_name("BADGER_PLUGIN_ROOT") - assert display_name == "plugin root" + ), + ): + config_singleton = ConfigSingleton(self.config_file) + display_name = config_singleton.read_display_name("BADGER_PLUGIN_ROOT") + assert display_name == "plugin root" def test_read_description(self): - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - description = config_singleton.read_description("BADGER_PLUGIN_ROOT") - assert description == "Mock plugin root" + ), + ): + config_singleton = ConfigSingleton(self.config_file) + description = config_singleton.read_description("BADGER_PLUGIN_ROOT") + assert description == "Mock plugin root" def test_read_is_path(self): - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - is_path = config_singleton.read_is_path("BADGER_PLUGIN_ROOT") - assert is_path + ), + ): + config_singleton = ConfigSingleton(self.config_file) + is_path = config_singleton.read_is_path("BADGER_PLUGIN_ROOT") + assert is_path def test_write_value(self): - with patch("os.path.exists", return_value=True): - with patch( + with ( + patch("os.path.exists", return_value=True), + patch( "builtins.open", mock_open( read_data=yaml.dump( self.mock_badger_config.model_dump(by_alias=True) ) ), - ): - config_singleton = ConfigSingleton(self.config_file) - config_singleton.write_value("BADGER_PLUGIN_ROOT", "/new/plugin/root") - assert ( - config_singleton.config.BADGER_PLUGIN_ROOT.value - == "/new/plugin/root" - ) + ), + ): + config_singleton = ConfigSingleton(self.config_file) + config_singleton.write_value("BADGER_PLUGIN_ROOT", "/new/plugin/root") + assert ( + config_singleton.config.BADGER_PLUGIN_ROOT.value == "/new/plugin/root" + ) def test_reset_settings(self): - with patch("os.path.exists", return_value=False): - with patch( - "badger.settings.BadgerConfig", return_value=self.mock_badger_config - ): - config_singleton = ConfigSingleton(self.config_file) - with patch.object( - config_singleton, "update_and_save_config" - ) as mock_update: - config_singleton.reset_settings() - assert mock_update.called - assert isinstance(config_singleton.config, BadgerConfig) + with ( + patch("os.path.exists", return_value=False), + patch("badger.settings.BadgerConfig", return_value=self.mock_badger_config), + ): + config_singleton = ConfigSingleton(self.config_file) + with patch.object( + config_singleton, "update_and_save_config" + ) as mock_update: + config_singleton.reset_settings() + assert mock_update.called + assert isinstance(config_singleton.config, BadgerConfig) # TODO: Missing test for mock_settings method diff --git a/src/badger/tests/utils.py b/src/badger/tests/utils.py index 60392438..f57b890d 100644 --- a/src/badger/tests/utils.py +++ b/src/badger/tests/utils.py @@ -1,8 +1,9 @@ import os + import pandas as pd from xopt import VOCS -from xopt.generators.bayesian import UpperConfidenceBoundGenerator from xopt.generators import RandomGenerator +from xopt.generators.bayesian import UpperConfidenceBoundGenerator def create_routine(): @@ -186,8 +187,8 @@ def get_vars_in_row(routine, idx=0): def fix_path_issues(): - from badger.settings import init_settings from badger.archive import BADGER_ARCHIVE_ROOT + from badger.settings import init_settings config_singleton = init_settings() BADGER_TEMPLATE_ROOT = config_singleton.read_value("BADGER_TEMPLATE_ROOT") diff --git a/src/badger/tests/x-test_db.py b/src/badger/tests/x_test_db.py similarity index 87% rename from src/badger/tests/x-test_db.py rename to src/badger/tests/x_test_db.py index 25ccfd23..98e3f948 100644 --- a/src/badger/tests/x-test_db.py +++ b/src/badger/tests/x_test_db.py @@ -1,6 +1,6 @@ class TestDB: def test_save_routine(self): - from badger.db import save_routine, remove_routine + from badger.db import remove_routine, save_routine from badger.tests.utils import create_routine, fix_db_path_issue fix_db_path_issue() @@ -11,7 +11,7 @@ def test_save_routine(self): remove_routine(routine.id) def test_load_routine(self): - from badger.db import save_routine, load_routine, remove_routine + from badger.db import load_routine, remove_routine, save_routine from badger.tests.utils import create_routine, fix_db_path_issue fix_db_path_issue() diff --git a/src/badger/tests/x-test_routine_id.py b/src/badger/tests/x_test_routine_id.py similarity index 93% rename from src/badger/tests/x-test_routine_id.py rename to src/badger/tests/x_test_routine_id.py index 943fd711..ba9bfa2e 100644 --- a/src/badger/tests/x-test_routine_id.py +++ b/src/badger/tests/x_test_routine_id.py @@ -1,5 +1,5 @@ def test_routines_same_name(): - from badger.db import save_routine, remove_routine + from badger.db import remove_routine, save_routine from badger.tests.utils import create_routine, fix_db_path_issue fix_db_path_issue() @@ -16,8 +16,8 @@ def test_routines_same_name(): def test_modify_routine_no_runs(qtbot): + from badger.db import list_routine, load_routine, remove_routine from badger.gui.components.routine_page import BadgerRoutinePage - from badger.db import list_routine, remove_routine, load_routine window = BadgerRoutinePage() qtbot.addWidget(window) @@ -48,8 +48,8 @@ def test_modify_routine_no_runs(qtbot): # TODO: write test for modifying name of routine with runs def test_modify_routine_name(qtbot): + from badger.db import list_routine, load_routine, remove_routine, save_run from badger.gui.components.routine_page import BadgerRoutinePage - from badger.db import list_routine, remove_routine, load_routine, save_run window = BadgerRoutinePage() qtbot.addWidget(window) @@ -88,8 +88,8 @@ def test_modify_routine_name(qtbot): # TODO: write test for modifying algorithm of routine with runs def test_modify_routine_algorithm(qtbot): + from badger.db import list_routine, load_routine, remove_routine, save_run from badger.gui.components.routine_page import BadgerRoutinePage - from badger.db import list_routine, remove_routine, load_routine, save_run window = BadgerRoutinePage() qtbot.addWidget(window) diff --git a/src/badger/utils.py b/src/badger/utils.py index bcc6d7da..e89ae964 100644 --- a/src/badger/utils.py +++ b/src/badger/utils.py @@ -7,10 +7,11 @@ import os import pathlib import sys -from datetime import datetime +from collections.abc import Iterable +from datetime import UTC, datetime from importlib import metadata from types import TracebackType -from typing import TYPE_CHECKING, Any, Iterable, Optional +from typing import TYPE_CHECKING, Any import yaml from PyQt5.QtWidgets import QLayout, QWidget @@ -46,9 +47,9 @@ def __enter__(self): def __exit__( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - exc_traceback: Optional[TracebackType], + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + exc_traceback: TracebackType | None, ): for widget in self.widgets: if not widget.signalsBlocked(): @@ -62,7 +63,7 @@ def __exit__( # https://github.com/yaml/pyyaml/issues/234#issuecomment-765894586 class Dumper(yaml.Dumper): def increase_indent(self, flow=False, indentless=False): - return super(Dumper, self).increase_indent(flow, False) + return super().increase_indent(flow, False) def get_yaml_string(content): @@ -164,26 +165,26 @@ def ts_to_str(ts: datetime, format: str = "lcls-log") -> str: def str_to_ts(timestr: str, format: str = "lcls-log") -> datetime: if format == "lcls-log": - return datetime.strptime(timestr, "%d-%b-%Y %H:%M:%S") + return datetime.strptime(timestr, "%d-%b-%Y %H:%M:%S").astimezone(UTC) elif format == "lcls-log-full": - return datetime.strptime(timestr, "%d-%b-%Y %H:%M:%S.%f") + return datetime.strptime(timestr, "%d-%b-%Y %H:%M:%S.%f").astimezone(UTC) elif format == "lcls-fname": - return datetime.strptime(timestr, "%Y-%m-%d-%H%M%S") + return datetime.strptime(timestr, "%Y-%m-%d-%H%M%S").astimezone(UTC) else: # ISO format return datetime.fromisoformat(timestr) def ts_float_to_str(ts_float: float, format: str = "lcls-log") -> str: - ts = datetime.fromtimestamp(ts_float) + ts = datetime.fromtimestamp(ts_float, tz=UTC) return ts_to_str(ts, format) def curr_ts() -> datetime: - return datetime.now() + return datetime.now(tz=UTC) def curr_ts_to_str(format: str = "lcls-log") -> str: - return ts_to_str(datetime.now(), format) + return ts_to_str(datetime.now(tz=UTC), format) def create_archive_run_filename(routine: "Routine", format: str = "lcls-fname") -> str: @@ -199,29 +200,35 @@ def create_archive_run_filename(routine: "Routine", format: str = "lcls-fname") return fname -def get_header(routine): +def get_header(routine: "Routine") -> list[str]: try: obj_names = routine.vocs.objective_names - except Exception: + except AttributeError: obj_names = [] try: var_names = routine.vocs.variable_names - except Exception: + except AttributeError: var_names = [] try: con_names = routine.vocs.constraint_names - except Exception: + except AttributeError: con_names = [] try: sta_names = routine.vocs.constant_names - except KeyError: + except AttributeError: sta_names = [] - return obj_names + con_names + var_names + sta_names + return list(obj_names) + list(con_names) + list(var_names) + list(sta_names) -def run_names_to_dict(run_names): - runs = {} +# FIX: Messy unclear function, should be refactored to be more clear and concise +def run_names_to_dict( + run_names: list[str], +) -> dict[str, dict[str, dict[str, list[str]]]]: + # Convert a list of run filenames to a nested dictionary structure organized by year, month, and day. + # Example output: + # "2026": {"2026-01": {"2026-01-15": ["run1.yaml", "run2.yaml"]}} + runs: dict[str, dict[str, dict[str, list[str]]]] = {} for name in run_names: name = os.path.basename( name @@ -233,19 +240,19 @@ def run_names_to_dict(run_names): try: year_dict = runs[year] - except Exception: + except KeyError: runs[year] = {} year_dict = runs[year] key_month = f"{year}-{month}" try: month_dict = year_dict[key_month] - except Exception: + except KeyError: year_dict[key_month] = {} month_dict = year_dict[key_month] key_day = f"{year}-{month}-{day}" try: day_list = month_dict[key_day] - except Exception: + except KeyError: month_dict[key_day] = [] day_list = month_dict[key_day] day_list.append(name) @@ -253,26 +260,26 @@ def run_names_to_dict(run_names): return runs -def convert_str_to_value(str): +def convert_str_to_value(s: str) -> str | int | float | bool: try: - return int(str) + return int(s) except ValueError: pass try: - return float(str) + return float(s) except ValueError: pass try: - return bool(str) + return bool(s) except ValueError: pass - return str + return s -def parse_rule(rule): +def parse_rule(rule: dict | str) -> dict: if type(rule) is str: return { "direction": rule, @@ -283,15 +290,15 @@ def parse_rule(rule): # rule is a dict try: direction = rule["direction"] - except Exception: + except KeyError: direction = "MINIMIZE" try: filter = rule["filter"] - except Exception: + except KeyError: filter = "ignore_nan" try: reducer = rule["reducer"] - except Exception: + except KeyError: reducer = "percentile_80" return { @@ -335,7 +342,7 @@ def state_to_dict(generator, data, include_data=True): # https://stackoverflow.com/a/18472142 -def strtobool(val): +def strtobool(val: str) -> bool: """Convert a string representation of truth to true (1) or false (0). True values are 'y', 'yes', 't', 'true', 'on', and '1'; false values are 'n', 'no', 'f', 'false', 'off', and '0'. Raises ValueError if @@ -351,7 +358,7 @@ def strtobool(val): elif val in ("n", "no", "f", "false", "off", "0"): return False else: - raise ValueError("invalid truth value %r" % (val,)) + raise ValueError(f"invalid truth value {val!r}") # https://stackoverflow.com/a/61901696/4263605