From 1df3264ddb878eaec18f160597a0525a5711b7c4 Mon Sep 17 00:00:00 2001 From: Davey Elder Date: Mon, 28 Sep 2026 13:50:52 -0400 Subject: [PATCH 1/3] Add test coverage for CLI, myopic, stochastic, and extensions framework Signed-off-by: Davey Elder --- tests/test_cli.py | 192 +++++++++++++ tests/test_evolution_updater.py | 17 ++ tests/test_framework_extension_helpers.py | 326 ++++++++++++++++++++++ tests/test_myopic_progress_mapper.py | 87 ++++++ tests/test_stochastic_sequencer.py | 53 ++++ 5 files changed, 675 insertions(+) create mode 100644 tests/test_evolution_updater.py create mode 100644 tests/test_framework_extension_helpers.py create mode 100644 tests/test_myopic_progress_mapper.py create mode 100644 tests/test_stochastic_sequencer.py diff --git a/tests/test_cli.py b/tests/test_cli.py index 5ba1fa1ee..73d0192dd 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 @@ -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 != 0 + 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 != 0 + 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 != 0 + 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 From 2d83cc26b51410ce72a6abfea74cda93fbe765ec Mon Sep 17 00:00:00 2001 From: Davey Elder Date: Mon, 28 Sep 2026 13:51:13 -0400 Subject: [PATCH 2/3] Stop refreshing databases for test collection Signed-off-by: Davey Elder --- tests/conftest.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) 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() From 1002c32f6c81e6930748093e40587c88f916a0ab Mon Sep 17 00:00:00 2001 From: Davey Elder Date: Mon, 28 Sep 2026 15:24:31 -0400 Subject: [PATCH 3/3] Specify exact expected exit codes in test_cli Signed-off-by: Davey Elder --- tests/test_cli.py | 20 ++++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/tests/test_cli.py b/tests/test_cli.py index 73d0192dd..6e1551b9a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -128,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() @@ -139,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 @@ -293,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 @@ -309,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) @@ -321,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 @@ -434,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) @@ -460,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) @@ -599,7 +599,7 @@ def test_cli_check_units_detects_issues(tmp_path: Path) -> None: args = ['check-units', str(INVALID_CURRENCY_DB), '--output', str(tmp_path)] result = runner.invoke(app, args, catch_exceptions=False) - assert result.exit_code != 0 + 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 @@ -613,7 +613,7 @@ def test_cli_check_units_detects_issues_silent(tmp_path: Path) -> None: args = ['check-units', str(INVALID_CURRENCY_DB), '--output', str(tmp_path), '--silent'] result = runner.invoke(app, args, catch_exceptions=False) - assert result.exit_code != 0 + 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 @@ -638,7 +638,7 @@ def test_cli_check_units_missing_database() -> None: args = ['check-units', 'non_existent_db.sqlite'] result = runner.invoke(app, args) - assert result.exit_code != 0 + assert result.exit_code == 2, 'missing database file should be rejected as a bad argument' assert 'non_existent_db.sqlite' in result.stderr