diff --git a/README.md b/README.md index 583035515..8d6df2363 100644 --- a/README.md +++ b/README.md @@ -178,7 +178,7 @@ config = TemoaConfig( time_sequencing='seasonal_timeslices', input_database='tutorial_database.sqlite', output_database='tutorial_database.sqlite', - solver_name='appsi_highs', + solver='appsi_highs', output_path=output_path, silent=False, ) @@ -238,7 +238,7 @@ scenario = "tutorial" scenario_mode = "perfect_foresight" input_database = "tutorial_database.sqlite" output_database = "tutorial_database.sqlite" -solver_name = "appsi_highs" +solver = "appsi_highs" ``` ### Configuration Options diff --git a/requirements-dev.txt b/requirements-dev.txt index 49a81cce6..33c98884a 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -1148,9 +1148,9 @@ pillow==12.3.0 \ --hash=sha256:fe3cca2e4e8a592be0f269a1ca4835c25199d9f3ce815c8491048f785b0a0198 \ --hash=sha256:ffd0c5368496f41b0944be820fcb7a838aa6e623d250b01acf2643939c3f99d7 # via matplotlib -pint==0.25.3 \ - --hash=sha256:27eb25143bd5de9fcc4d5a4b484f16faf6b4615aa93ece6b3373a8c1a3c1b97d \ - --hash=sha256:f8f5df6cf65314d74da1ade1bf96f8e3e4d0c41b51577ac53c49e7d44ca5acee +pint==0.26.1 \ + --hash=sha256:1bbde36eae57a5a289cd05081c6405618a5899814940752064ce351cd0204f71 \ + --hash=sha256:e982b129415c09c63308f314ae44697d83e5c96f253bcc0d5b833fc3329eb6d4 # via temoa platformdirs==4.11.7 \ --hash=sha256:4f41487eeeeeb07f3a6625e61d9bc0ae6809f92d3386dbd74392fbb76108104d \ diff --git a/requirements.txt b/requirements.txt index f4af38140..3da57ca78 100644 --- a/requirements.txt +++ b/requirements.txt @@ -557,9 +557,9 @@ pillow==12.3.0 \ --hash=sha256:fe3cca2e4e8a592be0f269a1ca4835c25199d9f3ce815c8491048f785b0a0198 \ --hash=sha256:ffd0c5368496f41b0944be820fcb7a838aa6e623d250b01acf2643939c3f99d7 # via matplotlib -pint==0.25.3 \ - --hash=sha256:27eb25143bd5de9fcc4d5a4b484f16faf6b4615aa93ece6b3373a8c1a3c1b97d \ - --hash=sha256:f8f5df6cf65314d74da1ade1bf96f8e3e4d0c41b51577ac53c49e7d44ca5acee +pint==0.26.1 \ + --hash=sha256:1bbde36eae57a5a289cd05081c6405618a5899814940752064ce351cd0204f71 \ + --hash=sha256:e982b129415c09c63308f314ae44697d83e5c96f253bcc0d5b833fc3329eb6d4 # via temoa platformdirs==4.11.7 \ --hash=sha256:4f41487eeeeeb07f3a6625e61d9bc0ae6809f92d3386dbd74392fbb76108104d \ diff --git a/temoa/_internal/run_actions.py b/temoa/_internal/run_actions.py index dcdc01bd5..d586dfa9a 100644 --- a/temoa/_internal/run_actions.py +++ b/temoa/_internal/run_actions.py @@ -3,12 +3,13 @@ """ import sqlite3 -from collections.abc import Generator, Iterable +from collections.abc import Generator, Iterable, Mapping from contextlib import contextmanager from logging import getLogger from pathlib import Path from sys import version_info from time import perf_counter +from typing import Any from pyomo.environ import ( Constraint, @@ -25,6 +26,7 @@ from temoa._internal.table_writer import TableWriter from temoa.core.config import TemoaConfig from temoa.core.model import TemoaModel +from temoa.core.solver_spec import DEFAULT_SOLVER_OPTIONS from temoa.data_processing.db_to_excel import make_excel logger = getLogger(__name__) @@ -174,6 +176,7 @@ def solve_instance( solver_name: str, silent: bool = False, solver_suffixes: Iterable[str] | None = None, + solver_options: Mapping[str, Any] | None = None, ) -> tuple[TemoaModel, SolverResults]: """ Solve the instance and return a loaded instance @@ -181,6 +184,8 @@ def solve_instance( 'duals' is supported in the Temoa Framework. Some solvers may not support duals. :param silent: Run silently :param solver_name: The name of the solver to request from the SolverFactory + :param solver_options: options to set on the solver (see resolve_solver_options). If None, + the Temoa defaults for the solver are used :param instance: the instance to solve :return: loaded instance """ @@ -202,28 +207,11 @@ def solve_instance( if solver_name == 'neos': raise NotImplementedError('Neos based solve is not currently supported') - # Solver Configuration - if solver_name == 'cbc': - pass - - elif solver_name == 'cplex': - # Note: these parameter values match mip-dev / PyPSA - # (see: https://pypsa-eur.readthedocs.io/en/latest/configuration.html) - optimizer.options['lpmethod'] = 4 # barrier - optimizer.options['solutiontype'] = 2 # non basic solution, ie no crossover - optimizer.options['barrier convergetol'] = 1.0e-3 - optimizer.options['feasopt tolerance'] = 1.0e-4 - - elif solver_name == 'gurobi': - # Note: these parameter values match mip-dev / PyPSA (see: https://pypsa-eur.readthedocs.io/en/latest/configuration.html) - optimizer.options['Method'] = 2 # barrier - optimizer.options['Crossover'] = 0 # non basic solution, ie no crossover - optimizer.options['BarConvTol'] = 1.0e-3 - optimizer.options['FeasibilityTol'] = 1.0e-4 - optimizer.options['BarOrder'] = -1 # auto ordering; 2-4x faster than AMD on large models - - elif solver_name == 'appsi_highs': - pass + # Solver Configuration (defaults live in temoa.core.solver_spec.DEFAULT_SOLVER_OPTIONS) + if solver_options is None: + solver_options = DEFAULT_SOLVER_OPTIONS.get(solver_name, {}) + for option, option_value in solver_options.items(): + optimizer.options[option] = option_value # Suffix Handling solver_suffixes_list: list[str] = [] diff --git a/temoa/_internal/temoa_sequencer.py b/temoa/_internal/temoa_sequencer.py index bd44044e3..d781aad1b 100644 --- a/temoa/_internal/temoa_sequencer.py +++ b/temoa/_internal/temoa_sequencer.py @@ -27,6 +27,7 @@ from temoa.core.config import TemoaConfig from temoa.core.model import TemoaModel from temoa.core.modes import TemoaMode +from temoa.core.solver_spec import resolve_solver_options from temoa.data_io.hybrid_loader import HybridLoader from temoa.extensions.method_of_morris.morris_sequencer import MorrisSequencer from temoa.extensions.modeling_to_generate_alternatives.mga_sequencer import MgaSequencer @@ -254,6 +255,7 @@ def _run_perfect_foresight(self) -> None: self.config.solver_name, silent=self.config.silent, solver_suffixes=suffixes, + solver_options=resolve_solver_options(self.config.solver), ) good_solve, msg = check_solve_status(self.pf_results) if not good_solve: diff --git a/temoa/core/config.py b/temoa/core/config.py index 3e68c9225..173b296b7 100644 --- a/temoa/core/config.py +++ b/temoa/core/config.py @@ -1,10 +1,14 @@ import shutil import sys import tomllib +import warnings +from collections.abc import Mapping from logging import getLogger from pathlib import Path +from typing import Any from temoa.core.modes import TemoaMode +from temoa.core.solver_spec import SolverSpec, redact_solver_options from temoa.extensions.framework import normalize_extension_ids, resolve_extension_specs logger = getLogger(__name__) @@ -42,7 +46,7 @@ def __init__( input_database: Path, output_database: Path, output_path: Path, - solver_name: str, + solver_name: str | None = None, neos: bool = False, save_excel: bool = False, save_duals: bool = False, @@ -73,6 +77,7 @@ def __init__( output_threshold_cost: float | None = None, sqlite: dict[str, object] | None = None, extensions: list[str] | tuple[str, ...] | None = None, + solver: str | Mapping[str, Any] | SolverSpec | None = None, ): if '-' in scenario: raise ValueError( @@ -129,7 +134,20 @@ def __init__( self.neos = neos if self.neos: raise NotImplementedError('Neos is currently not supported.') - self.solver_name = solver_name + + # Validate solver input + if solver_name is not None: + if solver is not None: + raise ValueError("Specify either 'solver' or 'solver_name', not both") + warnings.warn( + "The 'solver_name' argument is deprecated, use 'solver' instead", + DeprecationWarning, + stacklevel=2, + ) + solver = solver_name + if solver is None: + raise SolverNotAvailableError('No solver specified in the configuration.') + self.solver = SolverSpec.parse(solver) self.save_excel = save_excel self.save_duals = save_duals @@ -230,6 +248,17 @@ def __init__( if not self.silent: sys.stderr.write('Warning: ' + msg) + @property + def solver_name(self) -> str: + """The name of the selected solver (shorthand for self.solver.name)""" + return self.solver.name + + @solver_name.setter + def solver_name(self, value: str) -> None: + # retained for backward compatibility. Changing solvers drops any configured options, + # as they are specific to the previous solver + self.solver = SolverSpec.parse(value) + @staticmethod def _check_solver_availability(solver_name: str) -> tuple[bool, str | None]: """ @@ -282,27 +311,41 @@ def build_config(config_file: Path, output_path: Path, silent: bool = False) -> data = tomllib.load(f) if 'solver_name' in data: - is_available, location = TemoaConfig._check_solver_availability(data['solver_name']) - if not is_available: - error_message = ( - f"The specified solver '{data['solver_name']}' was not found.\n" - 'Please ensure the solver is installed and accessible.\n' + if 'solver' in data: + raise ValueError( + "Config specifies both 'solver' and 'solver_name'. Use only 'solver' " + "('solver_name' is deprecated)." ) - if data['solver_name'].lower() in SOLVER_DOC_LINKS: - link = SOLVER_DOC_LINKS[data['solver_name'].lower()] - error_message += f'For installation instructions, refer to: {link}\n' - else: - error_message += ( - "Refer to the solver's official documentation for " - 'installation instructions.' - ) - raise SolverNotAvailableError(error_message) + logger.warning( + "The 'solver_name' config key is deprecated and will be removed in a future " + 'release. Replace it with: solver = "%s"', + data['solver_name'], + ) + data['solver'] = data.pop('solver_name') + + if 'solver' not in data: + raise SolverNotAvailableError('No solver name specified in the configuration.') + solver = SolverSpec.parse(data['solver']) + data['solver'] = solver + + is_available, location = TemoaConfig._check_solver_availability(solver.name) + if not is_available: + error_message = ( + f"The specified solver '{solver.name}' was not found.\n" + 'Please ensure the solver is installed and accessible.\n' + ) + if solver.name.lower() in SOLVER_DOC_LINKS: + link = SOLVER_DOC_LINKS[solver.name.lower()] + error_message += f'For installation instructions, refer to: {link}\n' else: - logger.info('Using solver: %s (%s)', data['solver_name'], location) + error_message += ( + "Refer to the solver's official documentation for installation instructions." + ) + raise SolverNotAvailableError(error_message) else: - raise SolverNotAvailableError('No solver name specified in the configuration.') + logger.info('Using solver: %s (%s)', solver.name, location) - if data.get('solver_name') == 'appsi_highs' and data.get('save_duals', False): + if solver.name == 'appsi_highs' and data.get('save_duals', False): raise ValueError( 'save_duals is not supported with appsi_highs (it does not expose duals via the ' 'APPSI interface). Disable save_duals or choose a different solver.' @@ -354,6 +397,9 @@ def __repr__(self) -> str: msg += spacer msg += '{:>{}s}: {}\n'.format('Selected solver', width, self.solver_name) + msg += '{:>{}s}: {}\n'.format( + 'Solver options', width, redact_solver_options(self.solver.options) + ) msg += '{:>{}s}: {}\n'.format('NEOS status', width, self.neos) msg += spacer diff --git a/temoa/core/solver_spec.py b/temoa/core/solver_spec.py new file mode 100644 index 000000000..800810228 --- /dev/null +++ b/temoa/core/solver_spec.py @@ -0,0 +1,107 @@ +""" +Solver selection and option resolution. + +The config accepts ``solver`` as either a plain solver name or a table with a ``name`` and a +passthrough ``options`` table. Options that reach the solver are layered as: + + DEFAULT_SOLVER_OPTIONS < [solver.options] < extension-specific options (MGA, MC, ...) + +Any solver name known to pyomo's SolverFactory is accepted. Solvers without an entry in +DEFAULT_SOLVER_OPTIONS simply get no Temoa defaults. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Any + +# Note: these parameter values match mip-dev / PyPSA +# (see: https://pypsa-eur.readthedocs.io/en/latest/configuration.html) +DEFAULT_SOLVER_OPTIONS: dict[str, dict[str, Any]] = { + 'cplex': { + 'lpmethod': 4, # barrier + 'solutiontype': 2, # non basic solution, ie no crossover + 'barrier convergetol': 1.0e-3, + 'feasopt tolerance': 1.0e-4, + }, + 'gurobi': { + 'Method': 2, # barrier + 'Crossover': 0, # non basic solution, ie no crossover + 'BarConvTol': 1.0e-3, + 'FeasibilityTol': 1.0e-4, + 'BarOrder': -1, # auto ordering; 2-4x faster than AMD on large models + }, +} + + +@dataclass(frozen=True, slots=True) +class SolverSpec: + """A solver name plus the user-supplied options from the config (defaults excluded).""" + + name: str + options: Mapping[str, Any] = field(default_factory=dict) + + @classmethod + def parse(cls, raw: str | Mapping[str, Any] | SolverSpec) -> SolverSpec: + """ + Build a SolverSpec from a solver name, a {'name': ..., 'options': {...}} mapping, or an + existing SolverSpec + """ + if isinstance(raw, SolverSpec): + return raw + if isinstance(raw, str): + if not raw: + raise ValueError('Solver name must not be empty') + return cls(raw) + if isinstance(raw, Mapping): + unknown = set(raw) - {'name', 'options'} + if unknown: + raise ValueError( + f'Unrecognized key(s) in solver table: {sorted(unknown)}. Expected "name" ' + 'and (optionally) "options". Solver parameters belong under [solver.options]' + ) + name = raw.get('name') + if not isinstance(name, str) or not name: + raise ValueError('The solver table requires a non-empty "name" entry') + options = raw.get('options', {}) + if not isinstance(options, Mapping): + raise ValueError('The "options" entry of the solver table must be a table/dict') + return cls(name, dict(options)) + raise TypeError(f'solver must be a str or a table/dict, got: {type(raw).__name__}') + + +# substrings (lowercase) of option names whose values are credentials, e.g. gurobi's WLSSecret, +# CloudSecretKey, CSAPIAccessID, ServerPassword, LicenseID +_SENSITIVE_OPTION_MARKERS = ('secret', 'password', 'accessid', 'licenseid', 'key', 'token') + + +def redact_solver_options(options: Mapping[str, Any]) -> dict[str, Any]: + """ + Return a copy of the options that is safe to log or print, with credential values masked + """ + return { + option: '***' + if any(marker in option.lower() for marker in _SENSITIVE_OPTION_MARKERS) + else option_value + for option, option_value in options.items() + } + + +def resolve_solver_options( + spec: SolverSpec, + extension_options: Mapping[str, Any] | None = None, + *, + include_defaults: bool = True, +) -> dict[str, Any]: + """ + Merge the option layers that are passed to the solver + :param spec: the solver spec from the config + :param extension_options: options from an extension's own source (MGA/MC/Morris toml, + stochastic config), which take precedence over everything else + :param include_defaults: include DEFAULT_SOLVER_OPTIONS as the bottom layer. Solve paths that + historically ran without Temoa defaults (MGA base solve, stochastic) turn this off. + :return: a new dict of solver options + """ + defaults = DEFAULT_SOLVER_OPTIONS.get(spec.name, {}) if include_defaults else {} + return {**defaults, **spec.options, **(extension_options or {})} diff --git a/temoa/extensions/method_of_morris/morris.py b/temoa/extensions/method_of_morris/morris.py index 4ae1f308f..19dc114e9 100644 --- a/temoa/extensions/method_of_morris/morris.py +++ b/temoa/extensions/method_of_morris/morris.py @@ -16,6 +16,7 @@ from temoa._internal import run_actions from temoa._internal.table_writer import TableWriter from temoa.core.config import TemoaConfig +from temoa.core.solver_spec import resolve_solver_options from temoa.data_io.hybrid_loader import HybridLoader seed = 42 @@ -40,7 +41,11 @@ def evaluate( dp = DataPortal(data_dict={None: data}) instance = run_actions.build_instance(loaded_portal=dp, extensions=config.extensions) - mdl, res = run_actions.solve_instance(instance=instance, solver_name=config.solver_name) + mdl, res = run_actions.solve_instance( + instance=instance, + solver_name=config.solver_name, + solver_options=resolve_solver_options(config.solver), + ) status = run_actions.check_solve_status(res) if not status: raise RuntimeError('Bad solve during Method of Morris') diff --git a/temoa/extensions/method_of_morris/morris_evaluate.py b/temoa/extensions/method_of_morris/morris_evaluate.py index 9695b1406..daa3e8064 100644 --- a/temoa/extensions/method_of_morris/morris_evaluate.py +++ b/temoa/extensions/method_of_morris/morris_evaluate.py @@ -43,6 +43,7 @@ def evaluate( config: TemoaConfig, log_queue: Any, log_level: int, + solver_options: dict[str, Any] | None = None, ) -> list[float]: """ Run model for params provided and return objective value and emission value @@ -54,6 +55,7 @@ def evaluate( :param data: Data used to build the Data Portal :param i: indexing number :param config: The config file to pull run data from + :param solver_options: resolved solver options. If None, the Temoa defaults are used :return: list of objective value and CO2 emission value """ # get the logger configured... @@ -80,7 +82,10 @@ def evaluate( extensions=config.extensions, ) mdl, res = run_actions.solve_instance( - instance=instance, solver_name=config.solver_name, silent=True + instance=instance, + solver_name=config.solver_name, + silent=True, + solver_options=solver_options, ) status = run_actions.check_solve_status(res) if not status: diff --git a/temoa/extensions/method_of_morris/morris_sequencer.py b/temoa/extensions/method_of_morris/morris_sequencer.py index 6c3d1bc76..3a5baf194 100644 --- a/temoa/extensions/method_of_morris/morris_sequencer.py +++ b/temoa/extensions/method_of_morris/morris_sequencer.py @@ -22,6 +22,7 @@ from SALib.util import compute_groups_matrix, read_param_file # type: ignore[import-untyped] from temoa._internal.table_writer import TableWriter +from temoa.core.solver_spec import redact_solver_options, resolve_solver_options from temoa.data_io.hybrid_loader import HybridLoader from temoa.extensions.method_of_morris.morris_evaluate import evaluate @@ -70,11 +71,12 @@ def __init__(self, config: TemoaConfig): with open(path, 'rb') as f: all_options = tomllib.load(f) s_options = all_options.get(self.config.solver_name, {}) - logger.info('Using solver options: %s', s_options) except FileNotFoundError: logger.warning('Unable to find solver options toml file. Using default options.') s_options = {} + self.solver_options = resolve_solver_options(self.config.solver, s_options) + logger.info('Using solver options: %s', redact_solver_options(self.solver_options)) # output handling self.verbose = False # for troubleshooting @@ -189,7 +191,14 @@ def start(self) -> Any: sys.stdout.flush() morris_results = Parallel(n_jobs=self.num_cores)( delayed(evaluate)( - param_names, mm_samples[i, :], data, i, self.config, log_queue, log_level + param_names, + mm_samples[i, :], + data, + i, + self.config, + log_queue, + log_level, + solver_options=self.solver_options, ) for i in range(0, len(mm_samples)) ) diff --git a/temoa/extensions/modeling_to_generate_alternatives/MGA_solver_options.toml b/temoa/extensions/modeling_to_generate_alternatives/MGA_solver_options.toml index 869a53384..2b57d00c6 100644 --- a/temoa/extensions/modeling_to_generate_alternatives/MGA_solver_options.toml +++ b/temoa/extensions/modeling_to_generate_alternatives/MGA_solver_options.toml @@ -1,5 +1,7 @@ # A container for solver options # the top level solver name in brackets should align with the solver name in the config.toml +# These are layered on top of Temoa's solver defaults and the config's [solver.options], and take +# precedence over both. (see temoa/core/solver_spec.py) num_workers = 6 @@ -11,6 +13,7 @@ BarConvTol = 0.01 # Relative Barrier Tolerance primal-dual FeasibilityTol= 1e-2 # pretty loose Crossover= 0 # Disabled TimeLimit= 18000 # 5 hrs +BarOrder = -1 # auto ordering (gurobi default) # regarding BarConvTol: https://www.gurobi.com/documentation/current/refman/barrier_logging.html # note that ref above seems to imply that FeasibilyTol is NOT used when using barrier only...? @@ -19,6 +22,13 @@ TimeLimit= 18000 # 5 hrs # 'LogFile': './my_gurobi_log.log', # 'LPWarmStart': 2, # pass basis +[cplex] +# CPLEX's own defaults, restating them overrides the Temoa cplex defaults for MGA workers +lpmethod = 0 # automatic +solutiontype = 0 # automatic +'barrier convergetol' = 1.0e-8 +'feasopt tolerance' = 1.0e-6 + [cbc] # tbd diff --git a/temoa/extensions/modeling_to_generate_alternatives/mga_sequencer.py b/temoa/extensions/modeling_to_generate_alternatives/mga_sequencer.py index a261ba458..ebdfae220 100644 --- a/temoa/extensions/modeling_to_generate_alternatives/mga_sequencer.py +++ b/temoa/extensions/modeling_to_generate_alternatives/mga_sequencer.py @@ -33,6 +33,7 @@ from temoa._internal.run_actions import build_instance from temoa._internal.table_writer import TableWriter from temoa.components.costs import total_cost_rule +from temoa.core.solver_spec import redact_solver_options, resolve_solver_options from temoa.data_io.hybrid_loader import HybridLoader from temoa.extensions.modeling_to_generate_alternatives.manager_factory import get_manager from temoa.extensions.modeling_to_generate_alternatives.mga_constants import MgaAxis, MgaWeighting @@ -93,16 +94,23 @@ def __init__(self, config: TemoaConfig): with open(path, 'rb') as f: all_options = tomllib.load(f) s_options = all_options.get(self.config.solver_name, {}) - logger.info('Using solver options: %s', s_options) except FileNotFoundError: logger.warning('Unable to find solver options toml file. Using default options.') s_options = {} all_options = {} - # get handle on solver instance + # get handle on solver instance. The base solve only receives the user's + # [solver.options] (no Temoa defaults, no worker options) to get a more precise base cost self.opt = pyo.SolverFactory(self.config.solver_name) - self.worker_solver_options = s_options + for option, option_value in resolve_solver_options( + self.config.solver, include_defaults=False + ).items(): + self.opt.options[option] = option_value + self.worker_solver_options = resolve_solver_options(self.config.solver, s_options) + logger.info( + 'Using worker solver options: %s', redact_solver_options(self.worker_solver_options) + ) # some defaults, etc. self.internal_stop = False diff --git a/temoa/extensions/monte_carlo/MC_solver_options.toml b/temoa/extensions/monte_carlo/MC_solver_options.toml index da91cd56e..7a7471d61 100644 --- a/temoa/extensions/monte_carlo/MC_solver_options.toml +++ b/temoa/extensions/monte_carlo/MC_solver_options.toml @@ -1,5 +1,7 @@ # A container for solver options # the top level solver name in brackets should align with the solver name in the config.toml +# These are layered on top of Temoa's solver defaults and the config's [solver.options], and take +# precedence over both. (see temoa/core/solver_spec.py) num_workers = 11 @@ -11,6 +13,7 @@ BarConvTol = 1.0e-2 # Relative Barrier Tolerance primal-dual FeasibilityTol= 1.0e-2 # pretty loose Crossover= 0 # Disabled TimeLimit= 18000 # 5 hrs +BarOrder = -1 # auto ordering (gurobi default) # regarding BarConvTol: https://www.gurobi.com/documentation/current/refman/barrier_logging.html # note that ref above seems to imply that FeasibilyTol is NOT used when using barrier only...? @@ -19,6 +22,13 @@ TimeLimit= 18000 # 5 hrs # 'LogFile': './my_gurobi_log.log', # 'LPWarmStart': 2, # pass basis +[cplex] +# CPLEX's own defaults, restating them overrides the Temoa cplex defaults for MC workers +lpmethod = 0 # automatic +solutiontype = 0 # automatic +'barrier convergetol' = 1.0e-8 +'feasopt tolerance' = 1.0e-6 + [cbc] primalT = 1e-3 dualT = 1e-3 diff --git a/temoa/extensions/monte_carlo/mc_sequencer.py b/temoa/extensions/monte_carlo/mc_sequencer.py index eb99da89a..c5318e23f 100644 --- a/temoa/extensions/monte_carlo/mc_sequencer.py +++ b/temoa/extensions/monte_carlo/mc_sequencer.py @@ -18,6 +18,7 @@ from typing import TYPE_CHECKING, Any, cast from temoa._internal.table_writer import TableWriter +from temoa.core.solver_spec import redact_solver_options, resolve_solver_options from temoa.data_io.hybrid_loader import HybridLoader from temoa.extensions.monte_carlo.mc_run import MCRun, MCRunFactory from temoa.extensions.monte_carlo.mc_worker import MCWorker @@ -83,7 +84,6 @@ def __init__(self, config: TemoaConfig): with open(path, 'rb') as f: all_options = tomllib.load(f) s_options = all_options.get(self.config.solver_name, {}) - logger.info('Using solver options: %s', s_options) except FileNotFoundError: if options_file_path: @@ -95,7 +95,8 @@ def __init__(self, config: TemoaConfig): # worker options pulled from file self.num_workers = all_options.get('num_workers', 1) - self.worker_solver_options = s_options + self.worker_solver_options = resolve_solver_options(self.config.solver, s_options) + logger.info('Using solver options: %s', redact_solver_options(self.worker_solver_options)) # internal records self.solve_count = 0 diff --git a/temoa/extensions/myopic/myopic_sequencer.py b/temoa/extensions/myopic/myopic_sequencer.py index e32ed347f..8d01c0c04 100644 --- a/temoa/extensions/myopic/myopic_sequencer.py +++ b/temoa/extensions/myopic/myopic_sequencer.py @@ -16,6 +16,7 @@ from temoa._internal.table_writer import TableWriter from temoa.core.config import TemoaConfig from temoa.core.model import TemoaModel +from temoa.core.solver_spec import resolve_solver_options from temoa.data_io.hybrid_loader import HybridLoader from temoa.data_processing.db_to_excel import make_excel from temoa.extensions.myopic.myopic_index import MyopicIndex @@ -268,7 +269,10 @@ def start(self) -> None: if not self.config.silent and self.progress_mapper and idx: self.progress_mapper.report(idx, 'solve') model, results = run_actions.solve_instance( - instance=instance, solver_name=self.config.solver_name, silent=True + instance=instance, + solver_name=self.config.solver_name, + silent=True, + solver_options=resolve_solver_options(self.config.solver), ) optimal, status = run_actions.check_solve_status(results) diff --git a/temoa/extensions/single_vector_mga/sv_mga_sequencer.py b/temoa/extensions/single_vector_mga/sv_mga_sequencer.py index 543af6678..d0266c317 100644 --- a/temoa/extensions/single_vector_mga/sv_mga_sequencer.py +++ b/temoa/extensions/single_vector_mga/sv_mga_sequencer.py @@ -19,6 +19,7 @@ from temoa.components.costs import total_cost_rule from temoa.core.config import TemoaConfig from temoa.core.model import TemoaModel +from temoa.core.solver_spec import resolve_solver_options from temoa.data_io.hybrid_loader import HybridLoader from temoa.extensions.single_vector_mga.output_summary import summarize from temoa.model_checking.pricing_check import price_checker @@ -100,6 +101,7 @@ def start(self) -> None: solver_name=self.config.solver_name, silent=self.config.silent, solver_suffixes=suffixes, + solver_options=resolve_solver_options(self.config.solver), ) status = res.solver.termination_condition logger.debug('Termination condition: %s', status.name) @@ -163,6 +165,7 @@ def start(self) -> None: solver_name=self.config.solver_name, silent=self.config.silent, solver_suffixes=suffixes, + solver_options=resolve_solver_options(self.config.solver), ) status = res.solver.termination_condition logger.debug('Termination condition: %s', status.name) diff --git a/temoa/extensions/stochastics/stochastic_sequencer.py b/temoa/extensions/stochastics/stochastic_sequencer.py index 5c904ce30..195cf691a 100644 --- a/temoa/extensions/stochastics/stochastic_sequencer.py +++ b/temoa/extensions/stochastics/stochastic_sequencer.py @@ -5,6 +5,7 @@ import pyomo.environ as pyo +from temoa.core.solver_spec import resolve_solver_options from temoa.extensions.stochastics.stochastic_config import StochasticConfig if TYPE_CHECKING: @@ -50,8 +51,13 @@ def start(self) -> None: from temoa.extensions.stochastics.scenario_creator import scenario_creator - # Merge solver options from stoch_config - solver_options = self.stoch_config.solver_options.get(self.config.solver_name, {}) + # Merge solver options: [solver.options] < stoch_config. Temoa defaults are excluded to + # preserve the behavior of stochastic runs prior to the [solver] table + solver_options = resolve_solver_options( + self.config.solver, + self.stoch_config.solver_options.get(self.config.solver_name, {}), + include_defaults=False, + ) options = { 'solver': self.config.solver_name, diff --git a/temoa/model_checking/unit_checking/__init__.py b/temoa/model_checking/unit_checking/__init__.py index 0db3c443f..4cfcd2c1c 100644 --- a/temoa/model_checking/unit_checking/__init__.py +++ b/temoa/model_checking/unit_checking/__init__.py @@ -3,8 +3,7 @@ from pint import UnitRegistry from pint.errors import DefinitionSyntaxError -# UnitRegistry is generic but doesn't require type args at instantiation -ureg: UnitRegistry = UnitRegistry() # type: ignore[type-arg] +ureg: UnitRegistry = UnitRegistry() # Load custom unit definitions from the package resources _resource_path = 'temoa.model_checking.unit_checking/temoa_units.txt' diff --git a/temoa/tutorial_assets/config_sample.toml b/temoa/tutorial_assets/config_sample.toml index f903f330d..b262cf598 100644 --- a/temoa/tutorial_assets/config_sample.toml +++ b/temoa/tutorial_assets/config_sample.toml @@ -75,8 +75,15 @@ neos = false # solver (Mandatory) # Depending on what client machine has installed. -# [appsi_highs, cbc, gurobi, cplex, ...] -solver_name = "appsi_highs" +# [appsi_highs, cbc, gurobi, cplex, ...] (any solver available through pyomo's SolverFactory) +# Either a solver name, which uses Temoa's default options for that solver: +solver = "appsi_highs" +# or a name plus options passed through to the solver, which are merged over Temoa's defaults +# (use an inline table here, as a [solver] table header would capture the keys that follow it): +# solver = { name = "gurobi", options = { Method = 2, Crossover = 0, BarConvTol = 1.0e-3, FeasibilityTol = 1.0e-4, BarOrder = -1 } } +# Keep license credentials (e.g. gurobi WLS keys) in the solver's license file (gurobi.lic), not +# in these options. +# Note: 'solver_name' is deprecated but still accepted in place of 'solver' # ------------------------------------ # OUTPUTS diff --git a/tests/conftest.py b/tests/conftest.py index 444aebd50..7cb63eb9a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -185,8 +185,11 @@ def create_unit_test_dbs() -> None: logger.info('Created unit test DB: %s', db_name) -def pytest_configure(config: Config) -> None: # noqa: ARG001 +def pytest_configure(config: Config) -> None: """Setup test databases before test collection.""" + # Skip for collect-only (e.g. IDE test discovery) so it can't race a real run for DB locks. + if config.getoption('collectonly'): + return refresh_databases() try: create_unit_test_dbs() diff --git a/tests/test_cli.py b/tests/test_cli.py index 5ba1fa1ee..6e1551b9a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -4,6 +4,7 @@ from pathlib import Path import pytest +import tomlkit from typer.testing import CliRunner from temoa.cli import _is_writable, app @@ -127,7 +128,7 @@ def test_cli_validate_failure_on_invalid_db(tmp_path: Path) -> None: args = ['validate', str(test_config_path), '--output', str(tmp_path)] result = runner.invoke(app, args) - assert result.exit_code != 0, 'CLI should exit with a non-zero code on failure' + assert result.exit_code == 1, 'validate should exit 1 on a failed validation' assert 'Validation failed' in result.stdout # Check that the log was still created, containing the detailed error assert (tmp_path / 'temoa-run.log').exists() @@ -138,7 +139,7 @@ def test_cli_run_missing_config() -> None: args = ['run', 'non_existent_file.toml'] result = runner.invoke(app, args) - assert result.exit_code != 0 + assert result.exit_code == 2, 'missing config file should be rejected as a bad argument' # Check that the error mentions the missing file (more robust than exact string match) assert 'non_existent_file.toml' in result.stderr @@ -292,7 +293,7 @@ def mock_is_writable_always_false(_path: Path) -> bool: args = ['migrate', str(input_file)] result = runner.invoke(app, args, catch_exceptions=False) - assert result.exit_code != 0, 'Migration should fail with a non-zero exit code' + assert result.exit_code == 1, 'migrate should exit 1 when no writable output location exists' # Normalize whitespace to handle platform-specific line breaks from rich.print() normalized_output = ' '.join(result.stdout.split()) assert 'Error: Neither input directory' in normalized_output @@ -308,7 +309,7 @@ def test_cli_migrate_invalid_file() -> None: args = ['migrate', 'non_existent.sql'] result = runner.invoke(app, args) - assert result.exit_code != 0 + assert result.exit_code == 2, 'missing input file should be rejected as a bad argument' # Typer handles file existence check, so error is in stderr assert 'does not exist' in result.stderr or 'does not exist' in str(result.exception) @@ -320,7 +321,7 @@ def test_cli_migrate_unknown_type(tmp_path: Path) -> None: args = ['migrate', str(unknown_file)] result = runner.invoke(app, args) - assert result.exit_code != 0 + assert result.exit_code == 1, 'migrate should exit 1 for an undeterminable migration type' assert 'Cannot determine migration type' in result.stdout @@ -433,7 +434,7 @@ def test_cli_validate_fails_if_solver_missing( args = ['validate', str(test_config_path), '--output', str(tmp_path)] result = runner.invoke(app, args, catch_exceptions=False) - assert result.exit_code != 0, ( + assert result.exit_code == 1, ( f'Validate should have failed: {result.exception}\n{result.stderr}\n{result.stdout}' ) assert isinstance(result.exception, SystemExit) @@ -459,7 +460,7 @@ def test_cli_run_fails_if_solver_missing(tmp_path: Path, monkeypatch: pytest.Mon args = ['run', str(test_config_path), '--output', str(tmp_path)] result = runner.invoke(app, args, catch_exceptions=False) - assert result.exit_code != 0, ( + assert result.exit_code == 1, ( f'Run should have failed: {result.exception}\n{result.stderr}\n{result.stdout}' ) assert isinstance(result.exception, SystemExit) @@ -469,3 +470,194 @@ def test_cli_run_fails_if_solver_missing(tmp_path: Path, monkeypatch: pytest.Mon # Use the more robust phrase for checking installation instructions assert 'Please ensure the solver is installed and accessible.' in result.stdout assert (tmp_path / 'temoa-run.log').exists() + + +# ============================================================================= +# Tests for the `tutorial` command +# ============================================================================= + + +def test_cli_tutorial_creates_files(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """Test that `temoa tutorial` creates the config, database, and mc_settings files.""" + monkeypatch.chdir(tmp_path) + + args = ['tutorial', 'my_config', 'my_database'] + result = runner.invoke(app, args, catch_exceptions=False) + + assert result.exit_code == 0, f'CLI crashed with error: {result.exception}\n{result.stdout}' + assert (tmp_path / 'my_config.toml').exists() + assert (tmp_path / 'my_database.sqlite').exists() + assert (tmp_path / 'mc_settings.csv').exists() + assert 'Tutorial Setup Complete!' in result.stdout + + +def test_cli_tutorial_default_names(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """Test that `temoa tutorial` uses its default file names when none are given.""" + monkeypatch.chdir(tmp_path) + + result = runner.invoke(app, ['tutorial'], catch_exceptions=False) + + assert result.exit_code == 0, f'CLI crashed with error: {result.exception}\n{result.stdout}' + assert (tmp_path / 'tutorial_config.toml').exists() + assert (tmp_path / 'tutorial_database.sqlite').exists() + + +def test_cli_tutorial_updates_toml_database_paths( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Test that the generated config file points at the newly created database.""" + monkeypatch.chdir(tmp_path) + + result = runner.invoke(app, ['tutorial', 'cfg', 'db_name'], catch_exceptions=False) + + assert result.exit_code == 0, f'CLI crashed with error: {result.exception}\n{result.stdout}' + doc = tomlkit.parse((tmp_path / 'cfg.toml').read_text()) + assert doc['input_database'] == 'db_name.sqlite' + assert doc['output_database'] == 'db_name.sqlite' + + +def test_cli_tutorial_verbose_output(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """Test that `--verbose` prints the extra progress and guidance messages.""" + monkeypatch.chdir(tmp_path) + + result = runner.invoke(app, ['tutorial', 'cfg', 'db', '--verbose'], catch_exceptions=False) + + assert result.exit_code == 0 + assert 'Copying tutorial resources...' in result.stdout + assert 'Updating database paths in configuration...' in result.stdout + assert 'Tutorial files created successfully' in result.stdout + + +def test_cli_tutorial_existing_files_aborts_without_force( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Test that existing tutorial files trigger a confirmation prompt that can be declined.""" + monkeypatch.chdir(tmp_path) + (tmp_path / 'cfg.toml').write_text('placeholder') + + result = runner.invoke(app, ['tutorial', 'cfg', 'db'], input='n\n') + + # A declined confirmation is a graceful, non-error cancellation. + assert result.exit_code == 0 + assert 'Tutorial files already exist' in result.stdout + assert 'Tutorial setup cancelled' in result.stdout + # The placeholder file should be untouched since the user declined. + assert (tmp_path / 'cfg.toml').read_text() == 'placeholder' + + +def test_cli_tutorial_existing_files_force_overwrite( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Test that `--force` overwrites existing tutorial files without prompting.""" + monkeypatch.chdir(tmp_path) + (tmp_path / 'cfg.toml').write_text('placeholder') + + result = runner.invoke(app, ['tutorial', 'cfg', 'db', '--force'], catch_exceptions=False) + + assert result.exit_code == 0, f'CLI crashed with error: {result.exception}\n{result.stdout}' + assert (tmp_path / 'db.sqlite').exists() + assert (tmp_path / 'cfg.toml').read_text() != 'placeholder' + assert 'Tutorial setup cancelled' not in result.stdout + + +# ============================================================================= +# Tests for the `check-units` command +# ============================================================================= + +VALID_UNITS_DB = Path(__file__).parent / 'testing_outputs' / 'utopia_valid_units.sqlite' +INVALID_CURRENCY_DB = Path(__file__).parent / 'testing_outputs' / 'utopia_invalid_currency.sqlite' + +requires_unit_dbs = pytest.mark.skipif( + not (VALID_UNITS_DB.exists() and INVALID_CURRENCY_DB.exists()), + reason='Test databases not created. Ensure conftest.py setup completed successfully.', +) + + +@requires_unit_dbs +def test_cli_check_units_all_clear(tmp_path: Path) -> None: + """Test `temoa check-units` reports success on a valid database.""" + args = ['check-units', str(VALID_UNITS_DB), '--output', str(tmp_path)] + result = runner.invoke(app, args, catch_exceptions=False) + + assert result.exit_code == 0, f'CLI crashed with error: {result.exception}\n{result.stdout}' + assert 'All unit checks passed' in result.stdout + + +@requires_unit_dbs +def test_cli_check_units_all_clear_silent(tmp_path: Path) -> None: + """Test that `--silent` suppresses the success message.""" + args = ['check-units', str(VALID_UNITS_DB), '--output', str(tmp_path), '--silent'] + result = runner.invoke(app, args, catch_exceptions=False) + + assert result.exit_code == 0 + assert 'All unit checks passed' not in result.stdout + + +@requires_unit_dbs +def test_cli_check_units_detects_issues(tmp_path: Path) -> None: + """Test that `temoa check-units` fails and writes a report for a bad database.""" + args = ['check-units', str(INVALID_CURRENCY_DB), '--output', str(tmp_path)] + result = runner.invoke(app, args, catch_exceptions=False) + + assert result.exit_code == 1, 'check-units should exit 1 when issues are found' + assert 'Unit check found issues' in result.stdout + assert 'Detailed report saved to' in result.stdout + assert 'Report Summary:' in result.stdout + reports = list(tmp_path.glob('units_check_*.txt')) + assert len(reports) == 1 + + +@requires_unit_dbs +def test_cli_check_units_detects_issues_silent(tmp_path: Path) -> None: + """Test that `--silent` suppresses the issue report summary but still fails and writes it.""" + args = ['check-units', str(INVALID_CURRENCY_DB), '--output', str(tmp_path), '--silent'] + result = runner.invoke(app, args, catch_exceptions=False) + + assert result.exit_code == 1, 'check-units should exit 1 when issues are found' + assert 'Unit check found issues' not in result.stdout + reports = list(tmp_path.glob('units_check_*.txt')) + assert len(reports) == 1 + + +@requires_unit_dbs +def test_cli_check_units_default_output_dir( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Test that omitting `--output` defaults the report to ./unit_check_reports.""" + monkeypatch.chdir(tmp_path) + + args = ['check-units', str(VALID_UNITS_DB)] + result = runner.invoke(app, args, catch_exceptions=False) + + assert result.exit_code == 0 + assert (tmp_path / 'unit_check_reports').is_dir() + + +def test_cli_check_units_missing_database() -> None: + """Test graceful failure for a missing database file.""" + args = ['check-units', 'non_existent_db.sqlite'] + result = runner.invoke(app, args) + + assert result.exit_code == 2, 'missing database file should be rejected as a bad argument' + assert 'non_existent_db.sqlite' in result.stderr + + +# ============================================================================= +# Tests for `_is_writable` +# ============================================================================= + + +def test_is_writable_true_for_writable_dir(tmp_path: Path) -> None: + """Test that a normal, writable directory is reported as writable.""" + assert _is_writable(tmp_path) is True + + +def test_is_writable_false_on_oserror(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + """Test that `_is_writable` returns False when touching the probe file raises OSError.""" + + def _raise_oserror(*_args: object, **_kwargs: object) -> None: + raise OSError('mocked failure') + + monkeypatch.setattr(Path, 'touch', _raise_oserror) + + assert _is_writable(tmp_path) is False diff --git a/tests/test_evolution_updater.py b/tests/test_evolution_updater.py new file mode 100644 index 000000000..6d3efbaa9 --- /dev/null +++ b/tests/test_evolution_updater.py @@ -0,0 +1,17 @@ +"""Tests for the myopic evolution_updater template module.""" + +import logging + +import pytest + +from temoa.extensions.myopic.evolution_updater import iterate +from temoa.extensions.myopic.myopic_index import MyopicIndex + + +def test_iterate_logs_base_year(caplog: pytest.LogCaptureFixture) -> None: + idx = MyopicIndex(base_year=2020, step_year=2025, last_demand_year=2024, last_year=2030) + + with caplog.at_level(logging.INFO): + iterate(idx=idx, prev_base_year=2015, last_instance_status='optimal', db_con=None) + + assert 'base year 2020' in caplog.text diff --git a/tests/test_framework_extension_helpers.py b/tests/test_framework_extension_helpers.py new file mode 100644 index 000000000..4bc965322 --- /dev/null +++ b/tests/test_framework_extension_helpers.py @@ -0,0 +1,326 @@ +"""Tests for the extension-agnostic helper functions in temoa.extensions.framework. + +These cover the plumbing (id normalization, manifest/hook merging, and the +enabled/disabled extension table checks) that isn't exercised by +tests/test_extensions.py, which focuses on each concrete extension's model +components. +""" + +from __future__ import annotations + +import sqlite3 +from typing import TYPE_CHECKING, cast + +import pytest + +from temoa.extensions.framework import ( + ExtensionSpec, + _append_extension_schema, + _table_exists, + _table_has_rows, + append_extension_manifest_items, + apply_model_extension_hooks, + assert_disabled_extension_tables_are_empty, + ensure_enabled_extension_tables_exist, + get_known_extension_specs, + merge_regional_group_tables, + normalize_extension_ids, +) + +if TYPE_CHECKING: + from pathlib import Path + + from temoa.data_io.loader_manifest import LoadItem + + +# ============================================================================= +# normalize_extension_ids +# ============================================================================= + + +def test_normalize_extension_ids_none_returns_empty() -> None: + assert normalize_extension_ids(None) == () + + +def test_normalize_extension_ids_empty_list_returns_empty() -> None: + assert normalize_extension_ids([]) == () + + +def test_normalize_extension_ids_dedupes_and_lowercases_preserving_order() -> None: + result = normalize_extension_ids(['Growth_Rates', ' growth_rates ', 'discrete_capacity']) + assert result == ('growth_rates', 'discrete_capacity') + + +def test_normalize_extension_ids_skips_blank_entries() -> None: + assert normalize_extension_ids([' ', 'growth_rates']) == ('growth_rates',) + + +def test_normalize_extension_ids_rejects_non_string() -> None: + with pytest.raises(TypeError, match='Extension ids must be strings'): + normalize_extension_ids([123]) + + +# ============================================================================= +# merge_regional_group_tables +# ============================================================================= + + +def test_merge_regional_group_tables_merges_specs_into_base() -> None: + spec = ExtensionSpec(extension_id='ext_a', regional_group_tables={'tbl_a': 'field_a'}) + merged = merge_regional_group_tables({'tbl_base': 'field_base'}, [spec]) + assert merged == {'tbl_base': 'field_base', 'tbl_a': 'field_a'} + + +def test_merge_regional_group_tables_allows_identical_duplicate_mapping() -> None: + spec = ExtensionSpec(extension_id='ext_a', regional_group_tables={'tbl_base': 'field_base'}) + merged = merge_regional_group_tables({'tbl_base': 'field_base'}, [spec]) + assert merged == {'tbl_base': 'field_base'} + + +def test_merge_regional_group_tables_conflict_raises() -> None: + spec = ExtensionSpec(extension_id='ext_a', regional_group_tables={'tbl_base': 'field_other'}) + with pytest.raises(ValueError, match='conflicting field mappings'): + merge_regional_group_tables({'tbl_base': 'field_base'}, [spec]) + + +# ============================================================================= +# apply_model_extension_hooks / append_extension_manifest_items +# ============================================================================= + + +def test_apply_model_extension_hooks_calls_each_registered_hook() -> None: + calls: list[object] = [] + spec = ExtensionSpec(extension_id='ext_a', register_model_components=calls.append) + model = object() + + apply_model_extension_hooks(model, [spec]) # type: ignore[arg-type] + + assert calls == [model] + + +def test_apply_model_extension_hooks_skips_specs_without_hook() -> None: + spec = ExtensionSpec(extension_id='ext_a') + # Should not raise even though register_model_components is None. + apply_model_extension_hooks(object(), [spec]) # type: ignore[arg-type] + + +def test_append_extension_manifest_items_merges_in_order() -> None: + # Strings stand in for LoadItem; only list order matters here. + item_a = cast('LoadItem', 'item_a') + item_b = cast('LoadItem', 'item_b') + base_item = cast('LoadItem', 'base_item') + spec_a = ExtensionSpec(extension_id='ext_a', build_manifest_items=lambda _model: [item_a]) + spec_b = ExtensionSpec(extension_id='ext_b', build_manifest_items=lambda _model: [item_b]) + + merged = append_extension_manifest_items( + object(), # type: ignore[arg-type] + [base_item], + [spec_a, spec_b], + ) + + assert merged == [base_item, item_a, item_b] + + +# ============================================================================= +# _table_exists / _table_has_rows +# ============================================================================= + + +def test_table_exists_and_has_rows() -> None: + con = sqlite3.connect(':memory:') + try: + assert _table_exists(con, 'missing_table') is False + assert _table_has_rows(con, 'missing_table') is False + + con.execute('CREATE TABLE populated (id INTEGER)') + con.execute('CREATE TABLE empty_table (id INTEGER)') + con.execute('INSERT INTO populated VALUES (1)') + con.commit() + + assert _table_exists(con, 'populated') is True + assert _table_has_rows(con, 'populated') is True + assert _table_exists(con, 'empty_table') is True + assert _table_has_rows(con, 'empty_table') is False + finally: + con.close() + + +# ============================================================================= +# assert_disabled_extension_tables_are_empty +# ============================================================================= + + +def test_assert_disabled_extension_tables_are_empty_warns_when_populated( + caplog: pytest.LogCaptureFixture, +) -> None: + con = sqlite3.connect(':memory:') + try: + con.execute('CREATE TABLE limit_growth_capacity (region TEXT)') + con.execute("INSERT INTO limit_growth_capacity VALUES ('R1')") + con.commit() + + with caplog.at_level('WARNING'): + assert_disabled_extension_tables_are_empty(con, enabled_specs=()) + + assert any('growth_rates' in record.message for record in caplog.records) + finally: + con.close() + + +def test_assert_disabled_extension_tables_are_empty_silent_when_enabled( + caplog: pytest.LogCaptureFixture, +) -> None: + con = sqlite3.connect(':memory:') + try: + con.execute('CREATE TABLE limit_growth_capacity (region TEXT)') + con.execute("INSERT INTO limit_growth_capacity VALUES ('R1')") + con.commit() + + growth_rates_spec = get_known_extension_specs()['growth_rates'] + with caplog.at_level('WARNING'): + assert_disabled_extension_tables_are_empty(con, enabled_specs=(growth_rates_spec,)) + + assert not caplog.records + finally: + con.close() + + +# ============================================================================= +# ensure_enabled_extension_tables_exist / _append_extension_schema +# ============================================================================= + + +def test_ensure_enabled_extension_tables_exist_noop_when_tables_present() -> None: + con = sqlite3.connect(':memory:') + try: + con.execute('CREATE TABLE owned_table (id INTEGER)') + con.commit() + spec = ExtensionSpec(extension_id='ext_a', owned_tables=('owned_table',)) + + # Should not raise or prompt. + ensure_enabled_extension_tables_exist(con, [spec], input_database='db.sqlite', silent=True) + finally: + con.close() + + +def test_ensure_enabled_extension_tables_exist_no_schema_path_raises() -> None: + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec(extension_id='ext_a', owned_tables=('missing_table',)) + with pytest.raises(RuntimeError, match='No schema SQL path is registered'): + ensure_enabled_extension_tables_exist( + con, [spec], input_database='db.sqlite', silent=True + ) + finally: + con.close() + + +def test_ensure_enabled_extension_tables_exist_silent_skips_prompt_and_raises() -> None: + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec( + extension_id='ext_a', owned_tables=('missing_table',), schema_sql_path='unused.sql' + ) + # silent=True means the prompt is never asked, so should_apply stays False. + with pytest.raises(RuntimeError, match='Re-run and accept the prompt'): + ensure_enabled_extension_tables_exist( + con, [spec], input_database='db.sqlite', silent=True + ) + finally: + con.close() + + +def test_ensure_enabled_extension_tables_exist_prompt_declined_raises( + monkeypatch: pytest.MonkeyPatch, +) -> None: + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec( + extension_id='ext_a', owned_tables=('missing_table',), schema_sql_path='unused.sql' + ) + monkeypatch.setattr('builtins.input', lambda _prompt: 'n') + with pytest.raises(RuntimeError, match='Re-run and accept the prompt'): + ensure_enabled_extension_tables_exist( + con, [spec], input_database='db.sqlite', silent=False + ) + finally: + con.close() + + +def test_ensure_enabled_extension_tables_exist_prompt_accepted_applies_schema( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + schema_file = tmp_path / 'extra_schema.sql' + schema_file.write_text('CREATE TABLE missing_table (id INTEGER);') + + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec( + extension_id='ext_a', + owned_tables=('missing_table',), + schema_sql_path=str(schema_file), + ) + monkeypatch.setattr('builtins.input', lambda _prompt: 'y') + + ensure_enabled_extension_tables_exist(con, [spec], input_database='db.sqlite', silent=False) + + assert _table_exists(con, 'missing_table') is True + finally: + con.close() + + +def test_ensure_enabled_extension_tables_exist_still_missing_after_apply_raises( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + # Schema file exists but doesn't actually create the owned table. + schema_file = tmp_path / 'noop_schema.sql' + schema_file.write_text('CREATE TABLE unrelated_table (id INTEGER);') + + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec( + extension_id='ext_a', + owned_tables=('missing_table',), + schema_sql_path=str(schema_file), + ) + monkeypatch.setattr('builtins.input', lambda _prompt: 'y') + + with pytest.raises(RuntimeError, match='still missing'): + ensure_enabled_extension_tables_exist( + con, [spec], input_database='db.sqlite', silent=False + ) + finally: + con.close() + + +def test_append_extension_schema_no_path_raises() -> None: + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec(extension_id='ext_a') + with pytest.raises(RuntimeError, match='no schema SQL path configured'): + _append_extension_schema(con, spec) + finally: + con.close() + + +def test_append_extension_schema_missing_file_raises() -> None: + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec(extension_id='ext_a', schema_sql_path='/no/such/file.sql') + with pytest.raises(FileNotFoundError, match='not found'): + _append_extension_schema(con, spec) + finally: + con.close() + + +def test_append_extension_schema_executes_and_commits(tmp_path: Path) -> None: + schema_file = tmp_path / 'schema.sql' + schema_file.write_text('CREATE TABLE new_table (id INTEGER);') + + con = sqlite3.connect(':memory:') + try: + spec = ExtensionSpec(extension_id='ext_a', schema_sql_path=str(schema_file)) + _append_extension_schema(con, spec) + assert _table_exists(con, 'new_table') is True + finally: + con.close() diff --git a/tests/test_myopic_progress_mapper.py b/tests/test_myopic_progress_mapper.py new file mode 100644 index 000000000..db024d520 --- /dev/null +++ b/tests/test_myopic_progress_mapper.py @@ -0,0 +1,87 @@ +"""Tests for MyopicProgressMapper, the console progress visualizer for myopic solves.""" + +import re + +import pytest + +from temoa.extensions.myopic.myopic_index import MyopicIndex +from temoa.extensions.myopic.myopic_progress_mapper import MyopicProgressMapper + +YEARS = [2020, 2025, 2030, 2035] + + +def _index(base_year: int, step_year: int, last_demand_year: int) -> MyopicIndex: + return MyopicIndex( + base_year=base_year, + step_year=step_year, + last_demand_year=last_demand_year, + last_year=YEARS[-1] + 1, + ) + + +def test_init_computes_tag_width_and_positions() -> None: + mapper = MyopicProgressMapper(YEARS) + + assert mapper.years == YEARS + assert mapper.tag_width == max(len(str(y)) for y in YEARS) + 2 * len(mapper.leader) + # Positions are in increasing order, one per year. + assert list(mapper.pos.keys()) == YEARS + assert all(mapper.pos[YEARS[i]] < mapper.pos[YEARS[i + 1]] for i in range(len(YEARS) - 1)) + + +def test_draw_header_prints_years_and_label(capsys: pytest.CaptureFixture[str]) -> None: + mapper = MyopicProgressMapper(YEARS) + mapper.draw_header() + + out = capsys.readouterr().out + assert 'Myopic Progress' in out + assert 'HH:MM:SS' in out + for year in YEARS: + assert str(year) in out + + +def test_timestamp_format() -> None: + mapper = MyopicProgressMapper(YEARS) + assert re.match(r'^Elapsed: \d{2}:\d{2}:\d{2}\s+$', mapper.timestamp()) + + +@pytest.mark.parametrize( + 'status,tag', + [ + ('load', 'LOAD'), + ('solve', 'SOLV'), + ('check', 'CHEK'), + ('evolve', 'EVLV'), + ], +) +def test_report_prints_expected_tag_for_status( + capsys: pytest.CaptureFixture[str], status: str, tag: str +) -> None: + mapper = MyopicProgressMapper(YEARS) + idx = _index(base_year=2020, step_year=2025, last_demand_year=2025) + + mapper.report(idx, status) # type: ignore[arg-type] + + out = capsys.readouterr().out + # One tag per year from base_year through last_demand_year (2020, 2025). + assert out.count(tag) == 2 + assert 'Elapsed:' in out + + +def test_report_status_report_uses_step_year(capsys: pytest.CaptureFixture[str]) -> None: + mapper = MyopicProgressMapper(YEARS) + idx = _index(base_year=2020, step_year=2030, last_demand_year=2025) + + mapper.report(idx, 'report') + + out = capsys.readouterr().out + # One tag per year from base_year up to (not including) step_year: 2020, 2025. + assert out.count('RECD') == 2 + + +def test_report_rejects_invalid_status() -> None: + mapper = MyopicProgressMapper(YEARS) + idx = _index(base_year=2020, step_year=2025, last_demand_year=2025) + + with pytest.raises(ValueError, match='bad status'): + mapper.report(idx, 'bogus') # type: ignore[arg-type] diff --git a/tests/test_stochastic_sequencer.py b/tests/test_stochastic_sequencer.py new file mode 100644 index 000000000..84f506d2a --- /dev/null +++ b/tests/test_stochastic_sequencer.py @@ -0,0 +1,53 @@ +"""Tests for StochasticSequencer's constructor validation of stochastic config files.""" + +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from temoa.extensions.stochastics.stochastic_sequencer import StochasticSequencer + + +def _config(stochastic_config: Path | None) -> SimpleNamespace: + """A minimal stand-in for TemoaConfig with just the attribute the sequencer reads.""" + return SimpleNamespace(stochastic_config=stochastic_config) + + +def test_missing_stochastic_config_raises() -> None: + with pytest.raises(ValueError, match="requires a 'stochastic_config'"): + StochasticSequencer(_config(None)) # type: ignore[arg-type] + + +def test_nonexistent_stochastic_config_path_raises(tmp_path: Path) -> None: + missing = tmp_path / 'does_not_exist.toml' + with pytest.raises(ValueError, match='not found'): + StochasticSequencer(_config(missing)) # type: ignore[arg-type] + + +def test_stochastic_config_path_is_directory_raises(tmp_path: Path) -> None: + with pytest.raises(ValueError, match='is not a file'): + StochasticSequencer(_config(tmp_path)) # type: ignore[arg-type] + + +def test_invalid_toml_content_raises_wrapped_error(tmp_path: Path) -> None: + bad_toml = tmp_path / 'stoch.toml' + bad_toml.write_text('not valid = toml = content [[[') + + with pytest.raises(ValueError, match='Error parsing stochastic config'): + StochasticSequencer(_config(bad_toml)) # type: ignore[arg-type] + + +def test_valid_stochastic_config_loads_successfully(tmp_path: Path) -> None: + good_toml = tmp_path / 'stoch.toml' + good_toml.write_text( + """ + [scenarios] + base = 0.5 + high = 0.5 + """ + ) + + sequencer = StochasticSequencer(_config(good_toml)) # type: ignore[arg-type] + + assert sequencer.stoch_config.scenarios == {'base': 0.5, 'high': 0.5} + assert sequencer.objective_value is None diff --git a/uv.lock b/uv.lock index f1913064d..dbf2fca40 100644 --- a/uv.lock +++ b/uv.lock @@ -1474,7 +1474,7 @@ wheels = [ [[package]] name = "pint" -version = "0.25.3" +version = "0.26.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "flexcache" }, @@ -1482,9 +1482,9 @@ dependencies = [ { name = "platformdirs" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/52/9d/b1379cdbd33a49d17d627bc24e2b63cca06a1c5343b38072d2889499e82e/pint-0.25.3.tar.gz", hash = "sha256:f8f5df6cf65314d74da1ade1bf96f8e3e4d0c41b51577ac53c49e7d44ca5acee", size = 255106, upload-time = "2026-03-19T21:57:08.72Z" } +sdist = { url = "https://files.pythonhosted.org/packages/9f/bc/2c38c32e0fb1f966d3695f4493a3f3ce2cc0cca1bbe7c958b92261b31af9/pint-0.26.1.tar.gz", hash = "sha256:1bbde36eae57a5a289cd05081c6405618a5899814940752064ce351cd0204f71", size = 273631, upload-time = "2026-09-10T21:16:49.712Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/1b/dd/a9fe6a0a09512da23951c68bf36466aeecd89def3183dc095edbc807ddc5/pint-0.25.3-py3-none-any.whl", hash = "sha256:27eb25143bd5de9fcc4d5a4b484f16faf6b4615aa93ece6b3373a8c1a3c1b97d", size = 307488, upload-time = "2026-03-19T21:57:07.022Z" }, + { url = "https://files.pythonhosted.org/packages/6d/5c/9507ea3c732f8a259f45ddf190a741e9bf7d3fe229b5d473917ee6efbf1e/pint-0.26.1-py3-none-any.whl", hash = "sha256:e982b129415c09c63308f314ae44697d83e5c96f253bcc0d5b833fc3329eb6d4", size = 326478, upload-time = "2026-09-10T21:16:48.316Z" }, ] [[package]]