From b5b96520822c906d9eab362d1e25cc4c0070ff2f Mon Sep 17 00:00:00 2001 From: ShehaoLi Date: Sun, 17 May 2026 18:58:46 -0700 Subject: [PATCH] Harden sweep cleanup and review gaps --- README.md | 5 +- docs/operator_guide.md | 1 + docs/safety.md | 2 + pyproject.toml | 18 +-- requirements-dev.txt | 3 +- requirements.txt | 14 +-- src/app/application/errors.py | 8 ++ .../services/sweep/instrument_configurator.py | 2 +- .../services/sweep/waveform_acquirer.py | 2 + .../application/services/sweep_task_runner.py | 88 +++++++++----- src/app/application/use_cases/start_sweep.py | 5 +- src/app/domain/calibration.py | 6 +- .../infrastructure/instruments/awg_adapter.py | 5 +- .../infrastructure/instruments/osc_adapter.py | 5 +- src/equips.py | 14 +-- tests/test_architecture_boundaries.py | 23 +++- tests/test_auto_range_policy.py | 12 ++ tests/test_calibration.py | 53 +++++++++ tests/test_start_sweep_use_case.py | 73 ++++++++++-- tests/test_sweep_task_runner.py | 111 +++++++++++++++++- tests/test_validators.py | 52 ++++++++ 21 files changed, 418 insertions(+), 84 deletions(-) create mode 100644 tests/test_calibration.py create mode 100644 tests/test_validators.py diff --git a/README.md b/README.md index b89f29a..afa44c0 100644 --- a/README.md +++ b/README.md @@ -51,6 +51,7 @@ The UI and use cases do not call `src/equips.py` directly. That file is treated - For live instrument use: - supported AWG and oscilloscope models from `src/app/shared/mapping.py` - VISA access through `pyvisa` / `pyvisa-py` + - a working VISA backend for the connection type, such as NI-VISA / Keysight IO Libraries for LAN/USB/GPIB or the extra USB/GPIB libraries required by `pyvisa-py` - correct LAN/VISA addresses for the instruments Automated tests do not require AWG/OSC hardware. @@ -83,12 +84,14 @@ auto-load-off-test python src/main.py ``` -Settings are stored at: +Settings and auto-save data are rooted at the process working directory unless `AUTO_LOAD_OFF_TEST_ROOT` is set. From the repo root, settings are stored at: ```text __config__/settings.json ``` +For packaged installs or lab workstations, set `AUTO_LOAD_OFF_TEST_ROOT` to an explicit writable directory so settings and `__data__/measurement/` do not move when the app is launched from a different shell directory. + ## Run Tests Without Hardware ```bash diff --git a/docs/operator_guide.md b/docs/operator_guide.md index d892d3c..c64a3b6 100644 --- a/docs/operator_guide.md +++ b/docs/operator_guide.md @@ -56,6 +56,7 @@ The files in `demo_data/` can be loaded through the measurement loader path to i - No resources visible: check VISA backend, LAN connectivity, USB/GPIB cable, or serial permissions. - Sweep fails immediately: verify model label, address, impedance/coupling combinations, and numeric settings. +- Cleanup warning after Stop or window close: verify the AWG front-panel output state before touching the DUT or starting another sweep. - Flat or clipped waveform: reduce AWG amplitude or adjust oscilloscope range/offset. - Unexpected phase: verify reference channel, trigger mode, and cable/probe delays. - Save/load failure: confirm output directory permissions and supported file suffixes. diff --git a/docs/safety.md b/docs/safety.md index 1c99793..206e74b 100644 --- a/docs/safety.md +++ b/docs/safety.md @@ -19,7 +19,9 @@ This project is not a certified production test platform. It does not replace la - The sweep loop checks that event between frequency points and emits `SweepStopped` with the partial result. - Runner shutdown signals stop and waits briefly for the worker thread before closing instrument ports. - AWG shutdown attempts to turn the configured output channel off before closing the port. +- If the worker does not stop before the shutdown timeout, the runner emits `SHUTDOWN_TIMEOUT` and still attempts to turn AWG output off and close the known ports. - If output-off or port-close fails, the runner emits a `SweepWarning` so the UI/event log can surface the cleanup failure. +- Treat any cleanup warning after Stop or window close as hardware-significant: verify the AWG front panel/output indicator and the DUT state before touching the setup or starting another sweep. ## Exception Behavior diff --git a/pyproject.toml b/pyproject.toml index f67b2ff..65bc99f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,13 +9,13 @@ description = "Desktop AWG/oscilloscope sweep measurement automation tool." readme = "README.md" requires-python = ">=3.10" dependencies = [ - "numpy>=1.23", - "scipy>=1.10", - "matplotlib>=3.7", - "mplcursors>=0.5", - "pyvisa>=1.13", - "pyserial>=3.5", - "pyvisa-py>=0.7", + "numpy>=1.23,<3", + "scipy>=1.10,<2", + "matplotlib>=3.7,<4", + "mplcursors>=0.5,<1", + "pyvisa>=1.13,<2", + "pyserial>=3.5,<4", + "pyvisa-py>=0.7,<1", ] [project.scripts] @@ -23,10 +23,10 @@ auto-load-off-test = "main:main" [project.optional-dependencies] dev = [ - "ruff>=0.4", + "ruff>=0.4,<1", ] build = [ - "pyinstaller>=6", + "pyinstaller>=6,<7", ] [tool.setuptools] diff --git a/requirements-dev.txt b/requirements-dev.txt index 64394a9..0c870be 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -1,3 +1,2 @@ -r requirements.txt -ruff>=0.4 - +ruff>=0.4,<1 diff --git a/requirements.txt b/requirements.txt index 46bc67f..a9f5c95 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,7 +1,7 @@ -numpy>=1.23 -scipy>=1.10 -matplotlib>=3.7 -mplcursors>=0.5 -pyvisa>=1.13 -pyserial>=3.5 -pyvisa-py>=0.7 +numpy>=1.23,<3 +scipy>=1.10,<2 +matplotlib>=3.7,<4 +mplcursors>=0.5,<1 +pyvisa>=1.13,<2 +pyserial>=3.5,<4 +pyvisa-py>=0.7,<1 diff --git a/src/app/application/errors.py b/src/app/application/errors.py index 212a64b..e2042d5 100644 --- a/src/app/application/errors.py +++ b/src/app/application/errors.py @@ -15,3 +15,11 @@ class InstrumentAppError(ApplicationError): class PersistenceAppError(ApplicationError): pass + + +def describe_exception(exc: BaseException) -> str: + message = str(exc) + exc_type = type(exc).__name__ + if message: + return f"{exc_type}: {message}" + return exc_type diff --git a/src/app/application/services/sweep/instrument_configurator.py b/src/app/application/services/sweep/instrument_configurator.py index 45fe797..e6521db 100644 --- a/src/app/application/services/sweep/instrument_configurator.py +++ b/src/app/application/services/sweep/instrument_configurator.py @@ -21,7 +21,7 @@ def configure(self, settings: AppSettings) -> None: self._awg.reset() self._osc.reset() - self._awg.output_on(awg_ch) + self._awg.output_off(awg_ch) self._awg.set_impedance(setup.awg_settings.impedance.value, awg_ch) self._awg.set_amplitude_vpp(setup.awg_settings.amplitude_vpp, awg_ch) diff --git a/src/app/application/services/sweep/waveform_acquirer.py b/src/app/application/services/sweep/waveform_acquirer.py index f6b06d6..4b98935 100644 --- a/src/app/application/services/sweep/waveform_acquirer.py +++ b/src/app/application/services/sweep/waveform_acquirer.py @@ -29,6 +29,7 @@ def acquire(self, *, target_freq_hz: float, settings: AppSettings) -> AcquiredPo warnings: list[SweepServiceWarning] = [] awg_ch = setup.channels.awg_ch + self._awg.output_off(awg_ch) self._awg.set_frequency(float(target_freq_hz), awg_ch) actual_freq = self._awg.get_frequency(awg_ch) if not np.isclose(actual_freq, target_freq_hz, atol=1e-3, rtol=5e-6): @@ -58,6 +59,7 @@ def acquire(self, *, target_freq_hz: float, settings: AppSettings) -> AcquiredPo self._osc.set_timebase(window_s) triggered = run_mode.trigger_mode == TriggerMode.TRIGGERED + self._awg.output_on(awg_ch) self._osc.single_acquire(triggered=triggered) test_ch = setup.channels.osc_test_ch diff --git a/src/app/application/services/sweep_task_runner.py b/src/app/application/services/sweep_task_runner.py index 818a97b..1d16783 100644 --- a/src/app/application/services/sweep_task_runner.py +++ b/src/app/application/services/sweep_task_runner.py @@ -5,6 +5,7 @@ from pathlib import Path from app.application.dto import SaveTarget, StartSweepCommand +from app.application.errors import describe_exception from app.application.events import EventEmitter, SweepFailed, SweepWarning from app.application.ports.instruments import InstrumentPorts, InstrumentPortsFactory from app.application.use_cases.save_measurement import SaveMeasurementUseCase @@ -48,22 +49,36 @@ def start( if self.is_running(): return - ports = self._ports_factory(settings.setup) - stop_event = threading.Event() - self._stop_use_case = StopSweepUseCase(stop_event=stop_event) - - cmd = StartSweepCommand( - settings=settings, - calibration_enabled=calibration_enabled, - reference_interpolator=reference_interpolator, - ) - start_use_case = self._use_case_factory(awg=ports.awg, osc=ports.osc, stop_event=stop_event) - - with self._ports_lock: - self._ports = ports - self._active_awg_channel = settings.setup.channels.awg_ch - self._sweep_thread = threading.Thread(target=self._run_sweep, args=(start_use_case, cmd), daemon=True) - self._sweep_thread.start() + awg_channel = settings.setup.channels.awg_ch + ports: InstrumentPorts | None = None + try: + ports = self._ports_factory(settings.setup) + stop_event = threading.Event() + self._stop_use_case = StopSweepUseCase(stop_event=stop_event) + + cmd = StartSweepCommand( + settings=settings, + calibration_enabled=calibration_enabled, + reference_interpolator=reference_interpolator, + ) + start_use_case = self._use_case_factory(awg=ports.awg, osc=ports.osc, stop_event=stop_event) + + thread = threading.Thread(target=self._run_sweep, args=(start_use_case, cmd), daemon=True) + with self._ports_lock: + self._ports = ports + self._active_awg_channel = awg_channel + self._sweep_thread = thread + thread.start() + except Exception: + if ports is not None: + with self._ports_lock: + if self._ports is ports: + self._ports = None + self._active_awg_channel = None + self._close_port_set(ports=ports, awg_channel=awg_channel) + self._stop_use_case = None + self._sweep_thread = None + raise def stop(self) -> None: if self._stop_use_case is not None: @@ -79,26 +94,39 @@ def shutdown(self, timeout: float = 2.0) -> None: if self.is_running(): self._emit_warning( code="SHUTDOWN_TIMEOUT", - message="Sweep worker did not stop before shutdown timeout; ports will close when the worker exits.", + message=( + "Sweep worker did not stop before shutdown timeout; forcing AWG output off " + "and closing ports during shutdown." + ), ) + self._close_ports() return self._close_ports() def _run_sweep(self, start_use_case: StartSweepUseCase, cmd: StartSweepCommand) -> None: try: result = start_use_case.run(cmd, self._emitter) - if not result.is_empty and cmd.settings.auto_save_data: - target = SaveTarget( - base_path=self._auto_save_dir / "measurement", - include_timestamp=True, - figures={}, - ) - self._save_measurement_use_case.execute(result=result, settings=cmd.settings, target=target) except Exception as exc: # noqa: BLE001 - self._emitter.emit(SweepFailed(error_code="SWEEP_THREAD", message=str(exc))) + self._emitter.emit(SweepFailed(error_code="SWEEP_THREAD", message=describe_exception(exc))) + else: + self._auto_save_if_requested(result=result, cmd=cmd) finally: self._close_ports() + def _auto_save_if_requested(self, *, result, cmd: StartSweepCommand) -> None: + if result.is_empty or not cmd.settings.auto_save_data: + return + + target = SaveTarget( + base_path=self._auto_save_dir / "measurement", + include_timestamp=True, + figures={}, + ) + try: + self._save_measurement_use_case.execute(result=result, settings=cmd.settings, target=target) + except Exception as exc: # noqa: BLE001 + self._emit_warning(code="AUTO_SAVE_FAILED", message=describe_exception(exc)) + def _close_ports(self) -> None: with self._ports_lock: ports = self._ports @@ -109,22 +137,24 @@ def _close_ports(self) -> None: if ports is None: return + self._close_port_set(ports=ports, awg_channel=awg_channel) + + def _close_port_set(self, *, ports: InstrumentPorts, awg_channel: int | None) -> None: if awg_channel is not None: try: ports.awg.output_off(awg_channel) except Exception as exc: # noqa: BLE001 - self._emit_warning(code="AWG_OUTPUT_OFF_FAILED", message=str(exc)) + self._emit_warning(code="AWG_OUTPUT_OFF_FAILED", message=describe_exception(exc)) try: ports.awg.close() except Exception as exc: # noqa: BLE001 - self._emit_warning(code="AWG_CLOSE_FAILED", message=str(exc)) + self._emit_warning(code="AWG_CLOSE_FAILED", message=describe_exception(exc)) try: ports.osc.close() except Exception as exc: # noqa: BLE001 - self._emit_warning(code="OSC_CLOSE_FAILED", message=str(exc)) + self._emit_warning(code="OSC_CLOSE_FAILED", message=describe_exception(exc)) def _emit_warning(self, *, code: str, message: str) -> None: self._emitter.emit(SweepWarning(code=code, message=message)) - diff --git a/src/app/application/use_cases/start_sweep.py b/src/app/application/use_cases/start_sweep.py index b0fb1e8..d648bdb 100644 --- a/src/app/application/use_cases/start_sweep.py +++ b/src/app/application/use_cases/start_sweep.py @@ -5,6 +5,7 @@ from datetime import datetime, timezone from app.application.dto import StartSweepCommand +from app.application.errors import describe_exception from app.application.events import ( EventEmitter, SweepCompleted, @@ -98,8 +99,8 @@ def run(self, cmd: StartSweepCommand, emitter: EventEmitter) -> SweepResult: return result except ValidationError as exc: - emitter.emit(SweepFailed(error_code="VALIDATION", message=str(exc))) + emitter.emit(SweepFailed(error_code="VALIDATION", message=describe_exception(exc))) return SweepResult() except Exception as exc: # noqa: BLE001 - emitter.emit(SweepFailed(error_code="SWEEP_RUNTIME", message=str(exc))) + emitter.emit(SweepFailed(error_code="SWEEP_RUNTIME", message=describe_exception(exc))) return SweepResult() diff --git a/src/app/domain/calibration.py b/src/app/domain/calibration.py index 5445f26..d188b28 100644 --- a/src/app/domain/calibration.py +++ b/src/app/domain/calibration.py @@ -9,9 +9,9 @@ def build_reference_interpolator(curve: ReferenceCurve) -> Callable[[np.ndarray], np.ndarray]: - freq = np.asarray(curve.freq_hz, dtype=float).squeeze() - gain_db = np.asarray(curve.gain_db, dtype=float).squeeze() - phase = None if curve.phase_deg is None else np.asarray(curve.phase_deg, dtype=float).squeeze() + freq = np.atleast_1d(np.asarray(curve.freq_hz, dtype=float).squeeze()) + gain_db = np.atleast_1d(np.asarray(curve.gain_db, dtype=float).squeeze()) + phase = None if curve.phase_deg is None else np.atleast_1d(np.asarray(curve.phase_deg, dtype=float).squeeze()) if freq.size == 0: raise ValueError("Reference frequency data is empty") diff --git a/src/app/infrastructure/instruments/awg_adapter.py b/src/app/infrastructure/instruments/awg_adapter.py index 15d165b..dab792d 100644 --- a/src/app/infrastructure/instruments/awg_adapter.py +++ b/src/app/infrastructure/instruments/awg_adapter.py @@ -32,7 +32,4 @@ def get_amplitude_vpp(self, channel: int) -> float: return float(self._inst.get_amp(ch=channel)) def close(self) -> None: - try: - self._inst.inst_close() - except Exception: - pass + self._inst.inst_close() diff --git a/src/app/infrastructure/instruments/osc_adapter.py b/src/app/infrastructure/instruments/osc_adapter.py index 2422bb6..5a49850 100644 --- a/src/app/infrastructure/instruments/osc_adapter.py +++ b/src/app/infrastructure/instruments/osc_adapter.py @@ -54,7 +54,4 @@ def get_sample_rate(self) -> float: return float(self._inst.get_sample_rate()) def close(self) -> None: - try: - self._inst.inst_close() - except Exception: - pass + self._inst.inst_close() diff --git a/src/equips.py b/src/equips.py index 49b0773..64e6d2a 100644 --- a/src/equips.py +++ b/src/equips.py @@ -84,14 +84,12 @@ def inst_open(self): self.Inst = ResourceBase.open_VisaRM().open_resource(self.VisaAddress) return self.Inst - def inst_close(self): - if self.Inst: - try: - self.Inst.close() - except: - pass - finally: - self.Inst = None + def inst_close(self): + if self.Inst: + try: + self.Inst.close() + finally: + self.Inst = None def callback_after_open(self): pass diff --git a/tests/test_architecture_boundaries.py b/tests/test_architecture_boundaries.py index 2bf0710..f5dd0d4 100644 --- a/tests/test_architecture_boundaries.py +++ b/tests/test_architecture_boundaries.py @@ -29,7 +29,15 @@ def py_files(base: Path) -> list[Path]: class ArchitectureBoundaryTests(unittest.TestCase): def test_domain_stays_pure(self) -> None: - forbidden_prefixes = ("tkinter", "pyvisa", "serial", "matplotlib", "app.infrastructure", "app.presentation") + forbidden_prefixes = ( + "tkinter", + "pyvisa", + "serial", + "matplotlib", + "equips", + "app.infrastructure", + "app.presentation", + ) offenders = [] for path in py_files(SRC_APP / "domain"): for module in imported_modules(path): @@ -39,7 +47,14 @@ def test_domain_stays_pure(self) -> None: self.assertEqual(offenders, []) def test_application_does_not_import_infrastructure_or_presentation(self) -> None: - forbidden_prefixes = ("app.infrastructure", "app.presentation") + forbidden_prefixes = ( + "app.infrastructure", + "app.presentation", + "equips", + "pyvisa", + "serial", + "tkinter", + ) offenders = [] for path in py_files(SRC_APP / "application"): for module in imported_modules(path): @@ -49,10 +64,11 @@ def test_application_does_not_import_infrastructure_or_presentation(self) -> Non self.assertEqual(offenders, []) def test_presentation_does_not_import_infrastructure(self) -> None: + forbidden_prefixes = ("app.infrastructure", "equips", "pyvisa", "serial") offenders = [] for path in py_files(SRC_APP / "presentation"): for module in imported_modules(path): - if module.startswith("app.infrastructure"): + if module.startswith(forbidden_prefixes): offenders.append((path.relative_to(PROJECT_ROOT), module)) self.assertEqual(offenders, []) @@ -60,4 +76,3 @@ def test_presentation_does_not_import_infrastructure(self) -> None: if __name__ == "__main__": unittest.main() - diff --git a/tests/test_auto_range_policy.py b/tests/test_auto_range_policy.py index ab4cb24..3d1380c 100644 --- a/tests/test_auto_range_policy.py +++ b/tests/test_auto_range_policy.py @@ -49,6 +49,18 @@ def test_adjust_offset_when_midpoint_drifts(self) -> None: self.assertTrue(decision.changed) self.assertAlmostEqual(decision.target_offset_v, 0.45, delta=1e-6) + def test_empty_or_invalid_range_keeps_current_settings(self) -> None: + decision = self.policy.decide( + volts=np.array([]), + current_range_v=0.0, + current_offset_v=0.25, + requested_offset_v=0.0, + ) + + self.assertFalse(decision.changed) + self.assertEqual(decision.target_range_v, 0.0) + self.assertEqual(decision.target_offset_v, 0.25) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_calibration.py b/tests/test_calibration.py new file mode 100644 index 0000000..4c877fb --- /dev/null +++ b/tests/test_calibration.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +import sys +from pathlib import Path +import unittest + +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + +from app.domain.calibration import build_reference_interpolator +from app.domain.models import ReferenceCurve + + +class ReferenceInterpolatorTests(unittest.TestCase): + def test_single_point_reference_returns_constant(self) -> None: + interp = build_reference_interpolator( + ReferenceCurve(freq_hz=np.array([1_000.0]), gain_db=np.array([6.0]), phase_deg=None) + ) + + values = interp(np.array([100.0, 1_000.0, 10_000.0])) + + self.assertTrue(np.allclose(values, values[0])) + + def test_magnitude_reference_clamps_out_of_range(self) -> None: + interp = build_reference_interpolator( + ReferenceCurve(freq_hz=np.array([1_000.0, 2_000.0]), gain_db=np.array([0.0, 6.0]), phase_deg=None) + ) + + values = interp(np.array([100.0, 1_500.0, 5_000.0])) + + self.assertAlmostEqual(float(values[0]), 1.0, delta=1e-9) + self.assertAlmostEqual(float(values[-1]), 10 ** (6.0 / 20.0), delta=1e-9) + + def test_complex_reference_preserves_phase_and_clamps_edges(self) -> None: + interp = build_reference_interpolator( + ReferenceCurve( + freq_hz=np.array([1_000.0, 2_000.0, 3_000.0]), + gain_db=np.array([0.0, 0.0, 0.0]), + phase_deg=np.array([0.0, 45.0, 90.0]), + ) + ) + + values = interp(np.array([500.0, 2_000.0, 4_000.0])) + + self.assertTrue(np.iscomplexobj(values)) + self.assertAlmostEqual(float(np.angle(values[0], deg=True)), 0.0, delta=1e-9) + self.assertAlmostEqual(float(np.angle(values[1], deg=True)), 45.0, delta=1e-6) + self.assertAlmostEqual(float(np.angle(values[2], deg=True)), 90.0, delta=1e-9) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_start_sweep_use_case.py b/tests/test_start_sweep_use_case.py index a03b1b8..35bfaf1 100644 --- a/tests/test_start_sweep_use_case.py +++ b/tests/test_start_sweep_use_case.py @@ -10,7 +10,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) from app.application.dto import StartSweepCommand -from app.application.events import SweepCompleted, SweepProgress, SweepStarted, SweepStopped +from app.application.events import SweepCompleted, SweepFailed, SweepProgress, SweepStarted, SweepStopped from app.application.use_cases.start_sweep import StartSweepUseCase from app.domain.enums import ( ConnectionMode, @@ -36,33 +36,34 @@ class MockAwg: def __init__(self) -> None: self.freq = 1_000.0 self.amp = 1.0 + self.calls: list[tuple] = [] def reset(self) -> None: - return None + self.calls.append(("reset",)) def output_on(self, channel: int) -> None: - _ = channel + self.calls.append(("output_on", channel)) def output_off(self, channel: int) -> None: - _ = channel + self.calls.append(("output_off", channel)) def set_impedance(self, mode: str, channel: int) -> None: - _ = (mode, channel) + self.calls.append(("set_impedance", mode, channel)) def set_frequency(self, hz: float, channel: int) -> None: - _ = channel + self.calls.append(("set_frequency", hz, channel)) self.freq = hz def get_frequency(self, channel: int) -> float: - _ = channel + self.calls.append(("get_frequency", channel)) return self.freq def set_amplitude_vpp(self, vpp: float, channel: int) -> None: - _ = channel + self.calls.append(("set_amplitude_vpp", vpp, channel)) self.amp = vpp def get_amplitude_vpp(self, channel: int) -> float: - _ = channel + self.calls.append(("get_amplitude_vpp", channel)) return self.amp def close(self) -> None: @@ -191,6 +192,60 @@ def test_run_can_be_stopped(self) -> None: self.assertTrue(any(isinstance(e, SweepStopped) for e in recorder.events)) self.assertTrue(result.is_empty or len(result.points) >= 0) + def test_awg_output_waits_for_amplitude_and_frequency(self) -> None: + awg = MockAwg() + osc = MockOsc(awg) + stop_event = threading.Event() + + use_case = StartSweepUseCase(awg=awg, osc=osc, stop_event=stop_event) + recorder = Recorder() + + use_case.run(StartSweepCommand(settings=self._build_settings()), recorder) + + first_output_on = awg.calls.index(("output_on", 1)) + first_set_amp = awg.calls.index(("set_amplitude_vpp", 1.0, 1)) + first_set_freq = awg.calls.index(("set_frequency", 1000.0, 1)) + self.assertLess(first_set_amp, first_output_on) + self.assertLess(first_set_freq, first_output_on) + + def test_validation_failure_emits_failed_event(self) -> None: + settings = self._build_settings() + settings.setup.osc_settings.coupling = CouplingMode.AC + awg = MockAwg() + osc = MockOsc(awg) + use_case = StartSweepUseCase(awg=awg, osc=osc, stop_event=threading.Event()) + recorder = Recorder() + + result = use_case.run(StartSweepCommand(settings=settings), recorder) + + failures = [event for event in recorder.events if isinstance(event, SweepFailed)] + self.assertTrue(result.is_empty) + self.assertEqual(failures[0].error_code, "VALIDATION") + self.assertIn("ValidationError", failures[0].message) + + def test_runtime_failure_emits_exception_type(self) -> None: + class FailingConfigurator: + def configure(self, settings): + _ = settings + raise RuntimeError("configure failed") + + awg = MockAwg() + osc = MockOsc(awg) + use_case = StartSweepUseCase( + awg=awg, + osc=osc, + stop_event=threading.Event(), + configurator=FailingConfigurator(), + ) + recorder = Recorder() + + result = use_case.run(StartSweepCommand(settings=self._build_settings()), recorder) + + failures = [event for event in recorder.events if isinstance(event, SweepFailed)] + self.assertTrue(result.is_empty) + self.assertEqual(failures[0].error_code, "SWEEP_RUNTIME") + self.assertIn("RuntimeError", failures[0].message) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_sweep_task_runner.py b/tests/test_sweep_task_runner.py index 6926d7a..5d1244d 100644 --- a/tests/test_sweep_task_runner.py +++ b/tests/test_sweep_task_runner.py @@ -8,7 +8,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) -from app.application.events import SweepCompleted, SweepWarning +from app.application.events import SweepCompleted, SweepFailed, SweepWarning from app.domain.models import SweepPoint, SweepResult from app.infrastructure.persistence.settings_defaults import DefaultSettingsFactory from app.application.services.sweep_task_runner import SweepTaskRunner @@ -23,6 +23,12 @@ def execute(self, result, settings, target): return SimpleNamespace(mat_path=Path("measurement.mat")) +class FailingSaveMeasurementUseCase: + def execute(self, result, settings, target): + _ = (result, settings, target) + raise PermissionError("cannot write") + + class FakeEmitter: def __init__(self) -> None: self.events: list[object] = [] @@ -55,6 +61,16 @@ def run(self, cmd, emitter): return self._result +class BlockingStartUseCase: + def __init__(self, release: threading.Event) -> None: + self._release = release + + def run(self, cmd, emitter): + _ = (cmd, emitter) + self._release.wait(timeout=2.0) + return SweepResult() + + class SweepTaskRunnerTests(unittest.TestCase): def test_runner_auto_saves_completed_result(self) -> None: settings = DefaultSettingsFactory().create() @@ -119,6 +135,99 @@ def use_case_factory(*, awg, osc, stop_event: threading.Event): warnings = [event for event in emitter.events if isinstance(event, SweepWarning)] self.assertTrue(any(event.code == "AWG_CLOSE_FAILED" for event in warnings)) + def test_runner_warns_when_auto_save_fails_after_completed_sweep(self) -> None: + settings = DefaultSettingsFactory().create() + settings.auto_save_data = True + emitter = FakeEmitter() + awg = FakePort() + osc = FakePort() + result = SweepResult(points=[SweepPoint(freq_hz=1_000.0, gain_linear=1.0, gain_db=0.0)]) + + def ports_factory(_setup): + return SimpleNamespace(awg=awg, osc=osc, awg_address="A", osc_address="B") + + def use_case_factory(*, awg, osc, stop_event: threading.Event): + _ = (awg, osc, stop_event) + return FakeStartUseCase(result) + + runner = SweepTaskRunner( + emitter=emitter, + save_measurement_use_case=FailingSaveMeasurementUseCase(), + auto_save_dir=Path("."), + ports_factory=ports_factory, + use_case_factory=use_case_factory, + ) + + runner.start(settings=settings, calibration_enabled=False, reference_interpolator=None) + runner.wait(timeout=1.0) + + self.assertTrue(any(isinstance(event, SweepCompleted) for event in emitter.events)) + self.assertFalse(any(isinstance(event, SweepFailed) for event in emitter.events)) + warnings = [event for event in emitter.events if isinstance(event, SweepWarning)] + self.assertTrue(any(event.code == "AUTO_SAVE_FAILED" for event in warnings)) + + def test_start_cleans_up_ports_if_worker_setup_fails(self) -> None: + settings = DefaultSettingsFactory().create() + emitter = FakeEmitter() + awg = FakePort() + osc = FakePort() + + def ports_factory(_setup): + return SimpleNamespace(awg=awg, osc=osc, awg_address="A", osc_address="B") + + def use_case_factory(*, awg, osc, stop_event: threading.Event): + _ = (awg, osc, stop_event) + raise RuntimeError("factory failed") + + runner = SweepTaskRunner( + emitter=emitter, + save_measurement_use_case=FakeSaveMeasurementUseCase(), + auto_save_dir=Path("."), + ports_factory=ports_factory, + use_case_factory=use_case_factory, + ) + + with self.assertRaises(RuntimeError): + runner.start(settings=settings, calibration_enabled=False, reference_interpolator=None) + + self.assertEqual(awg.output_off_channels, [settings.setup.channels.awg_ch]) + self.assertTrue(awg.closed) + self.assertTrue(osc.closed) + self.assertFalse(runner.is_running()) + + def test_shutdown_timeout_forces_port_cleanup(self) -> None: + settings = DefaultSettingsFactory().create() + emitter = FakeEmitter() + awg = FakePort() + osc = FakePort() + release = threading.Event() + + def ports_factory(_setup): + return SimpleNamespace(awg=awg, osc=osc, awg_address="A", osc_address="B") + + def use_case_factory(*, awg, osc, stop_event: threading.Event): + _ = (awg, osc, stop_event) + return BlockingStartUseCase(release=release) + + runner = SweepTaskRunner( + emitter=emitter, + save_measurement_use_case=FakeSaveMeasurementUseCase(), + auto_save_dir=Path("."), + ports_factory=ports_factory, + use_case_factory=use_case_factory, + ) + + runner.start(settings=settings, calibration_enabled=False, reference_interpolator=None) + runner.shutdown(timeout=0.01) + release.set() + runner.wait(timeout=1.0) + + warnings = [event for event in emitter.events if isinstance(event, SweepWarning)] + self.assertTrue(any(event.code == "SHUTDOWN_TIMEOUT" for event in warnings)) + self.assertEqual(awg.output_off_channels, [settings.setup.channels.awg_ch]) + self.assertTrue(awg.closed) + self.assertTrue(osc.closed) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_validators.py b/tests/test_validators.py new file mode 100644 index 0000000..4f05ee2 --- /dev/null +++ b/tests/test_validators.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +import sys +from pathlib import Path +import unittest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + +from app.domain.enums import CouplingMode, CorrectionMode, ImpedanceMode, TriggerMode +from app.domain.models import ChannelSelection, OscSettings, SweepSpec +from app.domain.validators import ( + ValidationError, + validate_channels, + validate_osc_settings, + validate_sweep_spec, +) + + +class ValidatorTests(unittest.TestCase): + def test_log_sweep_requires_positive_step_count(self) -> None: + spec = SweepSpec(start_hz=1.0, stop_hz=10.0, step_hz=None, step_count=0, is_log=True) + + with self.assertRaises(ValidationError): + validate_sweep_spec(spec) + + def test_dual_correction_requires_reference_channel(self) -> None: + channels = ChannelSelection(awg_ch=1, osc_test_ch=1, osc_ref_ch=None, osc_trig_ch=2) + + with self.assertRaises(ValidationError): + validate_channels(channels, CorrectionMode.DUAL, TriggerMode.FREE_RUN) + + def test_triggered_mode_requires_trigger_channel(self) -> None: + channels = ChannelSelection(awg_ch=1, osc_test_ch=1, osc_ref_ch=2, osc_trig_ch=None) + + with self.assertRaises(ValidationError): + validate_channels(channels, CorrectionMode.NONE, TriggerMode.TRIGGERED) + + def test_50_ohm_ac_coupling_is_rejected(self) -> None: + settings = OscSettings( + full_scale_v=1.0, + offset_v=0.0, + points=1000, + impedance=ImpedanceMode.R50, + coupling=CouplingMode.AC, + ) + + with self.assertRaises(ValidationError): + validate_osc_settings(settings) + + +if __name__ == "__main__": + unittest.main()