diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..3a2a78a --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,19 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + branches: [main] + +jobs: + test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: "3.11" + - run: python -m pip install --upgrade pip + - run: python -m pip install -r requirements.txt + - run: PYTHONPATH=src python -m unittest discover -s tests diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..d4e7a21 --- /dev/null +++ b/.gitignore @@ -0,0 +1,27 @@ +.DS_Store + +# Python bytecode and test/tool caches +__pycache__/ +*.py[cod] +*$py.class +.pytest_cache/ +.mypy_cache/ +.ruff_cache/ +.matplotlib-cache/ +.coverage +htmlcov/ + +# Local environments +.venv/ +venv/ +env/ + +# Build and packaging output +build/ +dist/ +*.egg-info/ + +# Runtime files created by the desktop app +__config__/ +__data__/ +*.log diff --git a/README.md b/README.md index 7c497dd..b89f29a 100644 --- a/README.md +++ b/README.md @@ -1,68 +1,136 @@ -# LoadoffTest (Decoupled Architecture) +# Auto-Load-off-Test -LoadoffTest is a Python desktop tool for AWG/OSC sweep measurement and calibration. +[![CI](https://github.com/lishehao-ctrl/Auto-Load-off-Test/actions/workflows/ci.yml/badge.svg)](https://github.com/lishehao-ctrl/Auto-Load-off-Test/actions/workflows/ci.yml) -This version is fully refactored into a layered architecture: +Auto-Load-off-Test is a local Python desktop tool for AWG/oscilloscope sweep measurement, calibration, plotting, and data export. -- `presentation` (Tkinter UI only) -- `application` (use cases and event flow) -- `domain` (pure business models and algorithms) -- `infrastructure` (instrument adapters and persistence) +It turns a repetitive manual lab workflow into a layered application: -## Project Layout +- configure an arbitrary waveform generator (AWG) +- configure oscilloscope acquisition channels +- sweep frequency points +- measure gain and optional phase +- apply reference calibration +- export MAT/CSV/TXT data and optional plot images + +## Why It Exists + +Manual AWG/oscilloscope sweep measurements are repetitive and easy to misconfigure. This project separates the workflow into testable layers so the sweep math, signal processing, settings serialization, and use-case flow can be verified without physical instruments. + +## Architecture ```text src/ main.py app/ - presentation/tk/ - application/ - domain/ - infrastructure/ + bootstrap.py desktop composition root + runtime/ runtime paths and environment helpers + presentation/tk/ Tkinter UI and plotting + application/ use cases, DTOs, events, ports + domain/ pure models, validation, sweep math, DSP + infrastructure/ instrument adapters and persistence + equips.py legacy vendor/instrument compatibility layer +``` + +```mermaid +flowchart LR + UI["Tkinter UI"] --> APP["Application Use Cases"] + APP --> DOMAIN["Domain Models / Sweep / DSP"] + APP --> PORTS["Instrument Ports"] + PORTS --> INFRA["AWG / OSC Adapters"] + INFRA --> LEGACY["equips.py Vendor Layer"] + APP --> PERSIST["Settings + Measurement Persistence"] ``` -Legacy coupled modules (`src/ui.py`, `src/test.py`, `src/channel.py`, `src/deviceMng.py`) are removed. +The UI and use cases do not call `src/equips.py` directly. That file is treated as a legacy vendor compatibility layer and is wrapped by infrastructure adapters. -## Run +## Requirements + +- Python 3.10 or newer +- Tkinter, usually included with the Python installer on macOS/Windows +- For live instrument use: + - supported AWG and oscilloscope models from `src/app/shared/mapping.py` + - VISA access through `pyvisa` / `pyvisa-py` + - correct LAN/VISA addresses for the instruments + +Automated tests do not require AWG/OSC hardware. + +## Install ```bash -python3 src/main.py +python3 -m venv .venv +source .venv/bin/activate +python -m pip install --upgrade pip +python -m pip install -r requirements.txt ``` -## Configuration +For development tooling: -Settings are stored in JSON: +```bash +python -m pip install -r requirements-dev.txt +``` + +The project also exposes an optional console script when installed as a package: + +```bash +python -m pip install -e . +auto-load-off-test +``` + +## Run The Desktop App -- `__config__/settings.json` +```bash +python src/main.py +``` -Schema version is tracked in the settings payload (`schema_version`). +Settings are stored at: + +```text +__config__/settings.json +``` + +## Run Tests Without Hardware + +```bash +PYTHONPATH=src python -m unittest discover -s tests +``` + +The test suite uses pure domain tests and mocked instrument ports. It covers sweep generation, signal processing, settings serialization, measurement I/O, start-sweep event flow, and the sweep task runner. ## Output Files -Save operation writes: +Saving a measurement writes: - `*.mat` - `*.csv` - `*.txt` -- plot images (`*_gain.png`, `*_gain_db.png`) when figure handles are provided +- `*_gain.png` and `*_gain_db.png` when plot figures are supplied -## Testing +Auto-save writes timestamped files under: -Run automated tests: - -```bash -python3 -m unittest discover -s tests +```text +__data__/measurement/ ``` -Tests cover: +Example result generated from `demo_data/Demo(2).mat`: + +![Demo sweep result](docs/images/sweep_result.png) + +## Safety Notes + +This is a local lab automation tool, not a certified production test platform. Operators are responsible for confirming the connected instrument model, address, voltage range, frequency range, impedance, coupling, and device-under-test limits before running a live sweep. + +See [docs/safety.md](docs/safety.md) for stop/shutdown behavior and hardware assumptions. + +## Documentation -- domain sweep generation -- signal processing behavior -- start-sweep use case event flow with mock ports -- settings repository round-trip +- [Architecture](docs/architecture.md) +- [Operator Guide](docs/operator_guide.md) +- [Safety Notes](docs/safety.md) +- [Extending The Application](docs/extending.md) +- [Case Study](docs/case_study.md) +- [Demo Data](demo_data/README.md) -## Notes +## Project Status -- The application remains local single-process. -- No HTTP backend is introduced. -- UI thread safety is enforced through event queue dispatch (`Tk.after`). +The refactored app is local, single-process, and hardware-adapter based. Its strongest engineering signal is the separation between UI, use-case orchestration, pure domain logic, persistence, and instrument side effects. diff --git a/UserGuide/README.md b/UserGuide/README.md index 8b13789..337fcf6 100644 --- a/UserGuide/README.md +++ b/UserGuide/README.md @@ -1 +1,9 @@ +# User Guide +The original operator guide is kept as a Word document: + +- `网络分析仪_使用说明.docx` + +For GitHub review and day-to-day repository navigation, use the Markdown guide: + +- `../docs/operator_guide.md` diff --git a/demo_data/README.md b/demo_data/README.md index 8b13789..84c1871 100644 --- a/demo_data/README.md +++ b/demo_data/README.md @@ -1 +1,24 @@ +# Demo Data +This folder contains sample MAT files that can be used to inspect the measurement data shape without connecting instruments. + +## Files + +- `Deme(1).mat` + - Contains 15 frequency points. + - Keys observed: `freq`, `gain_db_raw`, `config`. + - Useful for checking older/raw gain-only measurement loading behavior. + +- `Demo(2).mat` + - Contains 50 frequency points. + - Keys observed: `freq`, `gain_db_corr`, `phase_corr`, `config`. + - Useful for checking corrected gain/phase measurement structure. + +## How To Use + +1. Start the desktop app with `python src/main.py`. +2. Use the load-measurement action. +3. Select one of the MAT files in this directory. +4. Confirm the plot and loaded point count look reasonable. + +These files are sample data for review and local testing. They are not a substitute for live instrument verification. diff --git a/docs/architecture.md b/docs/architecture.md index b6c631e..e18f105 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -1,17 +1,32 @@ # Architecture +Auto-Load-off-Test is organized as a local desktop application with explicit boundaries between UI code, use-case orchestration, pure domain logic, persistence, and hardware side effects. + +## Layer Diagram + +```mermaid +flowchart LR + UI["presentation/tk
Tkinter widgets, variables, dialogs, plots"] --> APP["application
use cases, DTOs, events, ports"] + APP --> DOMAIN["domain
models, validation, sweep math, DSP, calibration"] + APP --> PORTS["ports
AwgPort, OscPort, repositories"] + PORTS --> INFRA["infrastructure
adapters, scanner, JSON/MAT/CSV IO"] + INFRA --> LEGACY["src/equips.py
legacy vendor compatibility layer"] +``` + ## Layers +- `app/bootstrap.py` + - Desktop composition root. Wires repositories, use cases, scanner, instrument factories, runtime paths, and the Tk controller. - `app/presentation/tk` - Tk widgets, variable bindings, dialogs, chart rendering. - Consumes application events and dispatches user intents. - `app/application` - - Use-case orchestration (`start_sweep`, `stop_sweep`, `save/load`, `settings`). - - Emits typed events for UI; no Tk or message boxes. + - Use-case orchestration for start/stop sweep, save/load, reference loading, and settings. + - Emits typed events for UI; no Tk widgets or message boxes. - `app/domain` - - Pure dataclasses, enums, validation, sweep generation, DSP, calibration. + - Pure dataclasses, enums, validation, sweep generation, DSP, calibration, and export array shaping. - `app/infrastructure` - - Adapter wrappers around `equips.py`. + - Adapter wrappers around `src/equips.py`. - JSON settings and MAT/CSV/TXT persistence. ## Dependency Rules @@ -20,35 +35,45 @@ Allowed: - `presentation -> application` - `application -> domain` -- `application -> infrastructure` (ports / repositories) +- `application -> ports` +- `infrastructure -> ports` - `infrastructure -> domain` +- `infrastructure -> src/equips.py` Forbidden: -- `domain` importing Tkinter / PyVISA / Matplotlib -- `application` showing dialogs (`messagebox` / `filedialog`) -- UI accessing `equips` directly +- `domain` importing Tkinter, PyVISA, serial, or Matplotlib. +- `application` showing dialogs through `messagebox` or `filedialog`. +- UI or use cases accessing `src/equips.py` directly. ## Event Flow 1. UI collects parameters from `ViewModel`. -2. Controller maps to `AppSettings` and starts `StartSweepUseCase` in worker thread. -3. Use case emits: +2. `TkController` maps the view model to `AppSettings`. +3. `SweepTaskRunner` starts `StartSweepUseCase` in a worker thread. +4. Use case emits: - `SweepStarted` - `SweepProgress` - `SweepDataUpdated` - `SweepWarning` / `SweepFailed` - `SweepCompleted` / `SweepStopped` -4. Controller polls event queue on main thread via `after()` and updates UI safely. +5. Controller polls the event queue on the Tk main thread via `after()` and updates UI safely. ## Instrument Access -- Instrument model + address resolve through `equips_factory`. -- AWG and OSC commands are executed via `AwgPort` / `OscPort` adapters. +- Instrument model and address resolution go through `equips_factory`. +- AWG and OSC commands are executed through `AwgPort` and `OscPort` adapters. - Connection scanning is provided by `PyVisaResourceScanner` and `ConnectionMonitor`. +- `src/equips.py` is intentionally treated as a vendor compatibility layer. It contains legacy SCPI/serial behavior that should not be casually refactored without physical instrument verification. ## Persistence - Settings: `__config__/settings.json` -- Measurement files: MAT/CSV/TXT (+ optional plot PNG) +- Measurement files: MAT/CSV/TXT plus optional plot PNG files - Reference files: MAT + +Runtime locations are centralized through `AppPaths` in `app/runtime/paths.py`. + +## Test Strategy + +The automated tests stay hardware-free by using pure domain tests and fake instrument ports. Live instrument verification remains a manual/operator workflow. diff --git a/docs/case_study.md b/docs/case_study.md new file mode 100644 index 0000000..cacfd1f --- /dev/null +++ b/docs/case_study.md @@ -0,0 +1,43 @@ +# Case Study + +## Problem + +Manual AWG/oscilloscope sweep measurement is repetitive and error-prone. An operator must configure generator output, oscilloscope channels, trigger mode, acquisition timing, calibration/reference behavior, and data export for each run. + +## Constraints + +- The application controls physical instruments through VISA/LAN/serial paths. +- The UI must stay responsive while long sweeps run. +- Sweep math and signal processing should be testable without hardware. +- Instrument-specific commands should be isolated from application logic. +- Output data should be usable in analysis tools through MAT/CSV/TXT files. + +## Architecture + +The refactor separates the workflow into four main layers: + +- `presentation/tk`: Tkinter controls, dialogs, event handling, and plots. +- `application`: use cases, events, DTOs, and ports. +- `domain`: settings models, validation, sweep generation, signal processing, calibration, and export shaping. +- `infrastructure`: instrument adapters, resource scanning, settings persistence, and measurement IO. + +The legacy `src/equips.py` driver file remains as a vendor compatibility layer and is wrapped by infrastructure adapters. + +## Testing Strategy + +The automated tests avoid physical instruments by using: + +- pure tests for sweep generation, signal processing, auto range, and serialization +- fake AWG/OSC ports for the start-sweep use case +- temporary directories for measurement export/load round trips +- task-runner tests around threading, auto-save, cleanup, and warnings + +This keeps the core behavior reviewable on any development machine. + +## Output + +The app exports measurement data as MAT, CSV, and TXT files. Plot PNGs can be saved when the UI provides figure handles. + +## What This Demonstrates + +This project demonstrates real-world engineering in a physical-system context: separating hardware side effects from testable logic, preserving a practical desktop workflow, and improving maintainability without pretending the tool is a certified lab platform. diff --git a/docs/extending.md b/docs/extending.md new file mode 100644 index 0000000..2e214ec --- /dev/null +++ b/docs/extending.md @@ -0,0 +1,73 @@ +# Extending The Application + +This project is intentionally structured around extension seams rather than direct imports between every layer. + +## Composition Root + +`src/app/bootstrap.py` is the desktop composition root. It wires: + +- repositories +- use cases +- instrument scanner +- instrument port factory +- address resolver +- Tk controller +- runtime paths + +`src/main.py` should stay thin. If the app later gains CLI, scripted, or simulated run modes, add a new composition function instead of pushing more wiring into UI classes. + +## Runtime Paths + +`AppPaths` in `src/app/runtime/paths.py` centralizes the repo/runtime directories: + +- `__config__/settings.json` +- `__data__/` +- `__data__/measurement/` + +By default, runtime paths are rooted at the process working directory. Set `AUTO_LOAD_OFF_TEST_ROOT` to force a specific writable runtime location for packaged installs or lab workstations. + +Use `AppPaths` instead of recomputing `Path(__file__).parents[...]` in new code. + +## Adding A New Instrument + +1. Add or verify the model label in `src/app/shared/mapping.py`. +2. Add the vendor driver mapping in `src/equips.py` only if the low-level SCPI behavior is known. +3. Prefer adding behavior through `app.infrastructure.instruments` adapters rather than calling `equips.py` from UI or use cases. +4. Keep `AwgPort` / `OscPort` as the application contract. +5. Add hardware-free tests with fake ports before doing live bench validation. + +## Adding A New Persistence Format + +Keep use cases depending on repository ports. Format-specific logic belongs under `app.infrastructure.persistence`. + +For a new measurement format: + +1. Add a loader/exporter implementation in infrastructure. +2. Keep `SweepResult` and `AppSettings` as the domain boundary. +3. Add round-trip tests with temporary directories. +4. Do not put file dialogs or Tk concerns in persistence code. + +## Adding A New UI Field + +New user-facing sweep/settings fields usually touch these files: + +- `domain/models.py` +- `domain/validators.py` +- `infrastructure/persistence/settings_serializer.py` +- `infrastructure/persistence/settings_defaults.py` +- `presentation/tk/view_model.py` +- `presentation/tk/control_panel.py` +- `presentation/tk/mapper.py` +- tests for serializer and mapping behavior + +If a field changes instrument behavior, add coverage at the application service or use-case level with fake ports. + +## Boundary Tests + +`tests/test_architecture_boundaries.py` prevents the main layering rules from drifting: + +- domain stays pure +- application does not import presentation or infrastructure +- presentation does not import infrastructure + +Treat failures in those tests as design feedback, not just lint failures. diff --git a/docs/images/README.md b/docs/images/README.md new file mode 100644 index 0000000..3105f54 --- /dev/null +++ b/docs/images/README.md @@ -0,0 +1,8 @@ +# Screenshot Capture Notes + +Expected portfolio screenshots: + +- `main_ui.png`: the configured desktop app before a sweep. +- `sweep_result.png`: a completed or loaded sweep result. The current file is generated from `demo_data/Demo(2).mat`. + +Capture these from the real Tk desktop app. Do not replace them with generated mockups, because the value of this project is that it controls a real lab workflow. diff --git a/docs/images/sweep_result.png b/docs/images/sweep_result.png new file mode 100644 index 0000000..fff7681 Binary files /dev/null and b/docs/images/sweep_result.png differ diff --git a/docs/operator_guide.md b/docs/operator_guide.md new file mode 100644 index 0000000..d892d3c --- /dev/null +++ b/docs/operator_guide.md @@ -0,0 +1,70 @@ +# Operator Guide + +This guide summarizes the live workflow for the desktop app. The original Word guide in `UserGuide/` can remain as a detailed operator artifact, but this Markdown version is readable directly on GitHub. + +## 1. Connect Instruments + +1. Connect the AWG output to the device under test. +2. Connect the oscilloscope test channel to the measured output. +3. For dual-channel correction, connect the oscilloscope reference channel. +4. For triggered operation, connect or select the trigger channel. +5. Confirm VISA/LAN visibility with the instrument scanner or external VISA tooling. + +## 2. Configure Sweep Parameters + +- Start/stop frequency define the sweep range. +- Linear mode uses a frequency step. +- Log mode uses a step count. +- AWG amplitude is configured in Vpp. +- Oscilloscope range and offset define the vertical acquisition window. +- Coupling and impedance should match the probe, DUT, and measurement setup. + +## 3. Choose Correction And Trigger Mode + +- No correction: gain is computed against the configured AWG amplitude. +- Dual-channel correction: gain and phase are computed against a measured reference channel. +- Reference calibration: a loaded reference MAT file can correct measured points. +- Free-run mode captures without arming an edge trigger. +- Triggered mode arms the selected oscilloscope trigger channel. + +## 4. Run A Sweep + +1. Review hardware settings and safety limits. +2. Start the sweep. +3. Watch progress and warnings. +4. Stop if the DUT, waveform, range, or instrument state looks wrong. +5. Save the measurement if auto-save is disabled. + +## 5. Output Files + +Saving a measurement writes MAT, CSV, and TXT files. If plot figures are supplied, gain and gain-dB PNG files are also written. + +The CSV columns are: + +- `freq_hz` +- `gain_linear` +- `gain_db` +- `phase_deg` + +Missing phase values are exported as blank cells in CSV and `nan` values in TXT/MAT arrays. + +## 6. Demo Data + +The files in `demo_data/` can be loaded through the measurement loader path to inspect historical/sampled measurement structure without instruments. See `demo_data/README.md`. + +## 7. Troubleshooting + +- 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. +- 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. + +## 8. Screenshots + +For portfolio documentation, capture: + +- `docs/images/main_ui.png`: app configured before a sweep. +- `docs/images/sweep_result.png`: completed sweep with plotted result. + +Screenshots should be captured from the real desktop app rather than mocked or generated images. diff --git a/docs/safety.md b/docs/safety.md new file mode 100644 index 0000000..1c99793 --- /dev/null +++ b/docs/safety.md @@ -0,0 +1,44 @@ +# Safety Notes + +Auto-Load-off-Test controls physical lab instruments. It should be used as a local engineering tool by an operator who understands the connected AWG, oscilloscope, cables, probes, load impedance, and device-under-test limits. + +## Scope + +This project is not a certified production test platform. It does not replace lab safety procedures, instrument manuals, current/voltage limits, or operator judgment. + +## Hardware Assumptions + +- Supported model labels are defined in `src/app/shared/mapping.py`. +- Live operation uses VISA/LAN/serial access through `src/equips.py` via infrastructure adapters. +- Default settings are conservative examples, not a guarantee that a connected DUT is safe. +- The operator must verify AWG amplitude, frequency range, impedance, coupling mode, oscilloscope vertical range, and trigger configuration before starting a sweep. + +## Stop And Shutdown Behavior + +- Pressing Stop sets a shared stop event. +- 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 output-off or port-close fails, the runner emits a `SweepWarning` so the UI/event log can surface the cleanup failure. + +## Exception Behavior + +- Validation failures emit `SweepFailed` with a validation code. +- Runtime sweep failures emit `SweepFailed` and return an empty result. +- Cleanup failures should not hide the original sweep result, but they should be visible as warnings. + +## Operator Responsibility + +Before live measurement: + +1. Confirm the selected instrument models and VISA addresses. +2. Confirm load impedance and coupling. +3. Confirm AWG amplitude and sweep frequency limits. +4. Confirm oscilloscope range, offset, and trigger channel. +5. Keep physical access to instrument front panels and emergency stop procedures. + +Automated tests use mocked ports and do not validate real hardware behavior. + +## Runtime File Location + +Settings and auto-save output default to the process working directory. Set `AUTO_LOAD_OFF_TEST_ROOT` to use an explicit writable runtime directory on lab machines or packaged installs. diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..f67b2ff --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,42 @@ +[build-system] +requires = ["setuptools>=69", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "auto-load-off-test" +version = "0.1.0" +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", +] + +[project.scripts] +auto-load-off-test = "main:main" + +[project.optional-dependencies] +dev = [ + "ruff>=0.4", +] +build = [ + "pyinstaller>=6", +] + +[tool.setuptools] +package-dir = {"" = "src"} +py-modules = ["main", "equips", "mapping", "cvtTools"] + +[tool.setuptools.packages.find] +where = ["src"] +include = ["app*"] + +[tool.ruff] +line-length = 120 +target-version = "py310" diff --git a/requirements-dev.txt b/requirements-dev.txt new file mode 100644 index 0000000..64394a9 --- /dev/null +++ b/requirements-dev.txt @@ -0,0 +1,3 @@ +-r requirements.txt +ruff>=0.4 + diff --git a/requirements.txt b/requirements.txt index 8222286..46bc67f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,11 +1,7 @@ -requests -PyGithub -pyinstaller -numpy -scipy -matplotlib -mplcursors -pyvisa -pyserial -pyvisa-py -importlib-metadata; python_version < "3.8" +numpy>=1.23 +scipy>=1.10 +matplotlib>=3.7 +mplcursors>=0.5 +pyvisa>=1.13 +pyserial>=3.5 +pyvisa-py>=0.7 diff --git a/src/app/application/ports/__init__.py b/src/app/application/ports/__init__.py new file mode 100644 index 0000000..b4aa88b --- /dev/null +++ b/src/app/application/ports/__init__.py @@ -0,0 +1,11 @@ +from app.application.ports.instruments import AwgPort, OscPort, ResourceScannerPort +from app.application.ports.persistence import MeasurementRepository, ReferenceRepository, SettingsRepository + +__all__ = [ + "AwgPort", + "MeasurementRepository", + "OscPort", + "ReferenceRepository", + "ResourceScannerPort", + "SettingsRepository", +] diff --git a/src/app/application/ports/instruments.py b/src/app/application/ports/instruments.py new file mode 100644 index 0000000..1e44c2a --- /dev/null +++ b/src/app/application/ports/instruments.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Protocol + +import numpy as np + +from app.domain.models import InstrumentSetup + + +class AwgPort(Protocol): + def reset(self) -> None: ... + def output_on(self, channel: int) -> None: ... + def output_off(self, channel: int) -> None: ... + def set_impedance(self, mode: str, channel: int) -> None: ... + def set_frequency(self, hz: float, channel: int) -> None: ... + def get_frequency(self, channel: int) -> float: ... + def set_amplitude_vpp(self, vpp: float, channel: int) -> None: ... + def get_amplitude_vpp(self, channel: int) -> float: ... + def close(self) -> None: ... + + +class OscPort(Protocol): + def reset(self) -> None: ... + def output_on(self, channel: int) -> None: ... + def set_timebase(self, window_s: float, offset_s: float | None = None) -> None: ... + def set_vertical(self, channel: int, full_scale_v: float, offset_v: float) -> None: ... + def get_vertical(self, channel: int) -> tuple[float, float]: ... + def set_coupling(self, channel: int, mode: str) -> None: ... + def set_impedance(self, channel: int, mode: str) -> None: ... + def arm_trigger(self, channel: int, level_v: float) -> None: ... + def set_free_run(self) -> None: ... + def single_acquire(self, triggered: bool) -> None: ... + def read_waveform(self, channel: int, points: int | None) -> tuple[np.ndarray, np.ndarray]: ... + def get_sample_rate(self) -> float: ... + def close(self) -> None: ... + + +class ResourceScannerPort(Protocol): + def list_resources(self) -> tuple[str, ...]: ... + + +@dataclass(slots=True) +class InstrumentPorts: + awg: AwgPort + osc: OscPort + awg_address: str + osc_address: str + + +InstrumentPortsFactory = Callable[[InstrumentSetup], InstrumentPorts] diff --git a/src/app/application/ports/persistence.py b/src/app/application/ports/persistence.py new file mode 100644 index 0000000..7eaa96f --- /dev/null +++ b/src/app/application/ports/persistence.py @@ -0,0 +1,20 @@ +from __future__ import annotations + +from typing import Protocol + +from app.application.dto import LoadedMeasurement, SaveArtifacts, SaveTarget +from app.domain.models import AppSettings, ReferenceCurve, SweepResult + + +class SettingsRepository(Protocol): + def load(self) -> AppSettings: ... + def save(self, settings: AppSettings) -> None: ... + + +class MeasurementRepository(Protocol): + def save(self, result: SweepResult, settings: AppSettings, target: SaveTarget) -> SaveArtifacts: ... + def load(self, file_path: str) -> LoadedMeasurement: ... + + +class ReferenceRepository(Protocol): + def load_reference(self, file_path: str) -> ReferenceCurve: ... diff --git a/src/app/application/services/connection_monitor.py b/src/app/application/services/connection_monitor.py index 0aa901c..7997d18 100644 --- a/src/app/application/services/connection_monitor.py +++ b/src/app/application/services/connection_monitor.py @@ -5,7 +5,7 @@ from collections.abc import Callable from app.application.events import ConnectionStatusUpdated, EventEmitter -from app.infrastructure.instruments.ports import ResourceScannerPort +from app.application.ports.instruments import ResourceScannerPort class ConnectionMonitor: @@ -34,6 +34,8 @@ def start(self) -> None: def stop(self) -> None: self._stop.set() + if self._thread and self._thread.is_alive() and threading.current_thread() is not self._thread: + self._thread.join(timeout=self._interval_s * 2) def _run(self) -> None: while not self._stop.is_set(): @@ -41,8 +43,10 @@ def _run(self) -> None: osc_connected = False try: resources = self._scanner.list_resources() - awg_connected = self._get_awg_address() in resources if self._get_awg_address() else False - osc_connected = self._get_osc_address() in resources if self._get_osc_address() else False + awg_address = self._get_awg_address() + osc_address = self._get_osc_address() + awg_connected = awg_address in resources if awg_address else False + osc_connected = osc_address in resources if osc_address else False except Exception: awg_connected = False osc_connected = False diff --git a/src/app/application/services/sweep/__init__.py b/src/app/application/services/sweep/__init__.py new file mode 100644 index 0000000..472b809 --- /dev/null +++ b/src/app/application/services/sweep/__init__.py @@ -0,0 +1,17 @@ +from app.application.services.sweep.calibration_applier import CalibrationApplier +from app.application.services.sweep.instrument_configurator import InstrumentConfigurator +from app.application.services.sweep.models import AcquiredPointData, SweepPlan, SweepServiceWarning +from app.application.services.sweep.planner import SweepPlanner +from app.application.services.sweep.point_measurement_service import PointMeasurementService +from app.application.services.sweep.waveform_acquirer import WaveformAcquirer + +__all__ = [ + "AcquiredPointData", + "CalibrationApplier", + "InstrumentConfigurator", + "PointMeasurementService", + "SweepPlan", + "SweepPlanner", + "SweepServiceWarning", + "WaveformAcquirer", +] diff --git a/src/app/application/services/sweep/calibration_applier.py b/src/app/application/services/sweep/calibration_applier.py new file mode 100644 index 0000000..89f0a12 --- /dev/null +++ b/src/app/application/services/sweep/calibration_applier.py @@ -0,0 +1,19 @@ +from __future__ import annotations + +import numpy as np + +from app.application.dto import StartSweepCommand +from app.domain.calibration import apply_reference_to_point +from app.domain.enums import CorrectionMode, TriggerMode +from app.domain.models import SweepPoint + + +class CalibrationApplier: + def apply(self, *, point: SweepPoint, cmd: StartSweepCommand) -> SweepPoint: + if not cmd.calibration_enabled or cmd.reference_interpolator is None: + return point + + run_mode = cmd.settings.run_mode + use_phase = run_mode.correction_mode == CorrectionMode.DUAL or run_mode.trigger_mode == TriggerMode.TRIGGERED + ref_value = cmd.reference_interpolator(np.array([point.freq_hz]))[0] + return apply_reference_to_point(point, ref_value, use_phase=use_phase) diff --git a/src/app/application/services/sweep/instrument_configurator.py b/src/app/application/services/sweep/instrument_configurator.py new file mode 100644 index 0000000..45fe797 --- /dev/null +++ b/src/app/application/services/sweep/instrument_configurator.py @@ -0,0 +1,45 @@ +from __future__ import annotations + +from app.application.ports.instruments import AwgPort, OscPort +from app.domain.enums import CorrectionMode, TriggerMode +from app.domain.models import AppSettings + + +class InstrumentConfigurator: + def __init__(self, awg: AwgPort, osc: OscPort) -> None: + self._awg = awg + self._osc = osc + + def configure(self, settings: AppSettings) -> None: + setup = settings.setup + run_mode = settings.run_mode + + awg_ch = setup.channels.awg_ch + test_ch = setup.channels.osc_test_ch + + if run_mode.auto_reset: + self._awg.reset() + self._osc.reset() + + self._awg.output_on(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) + + self._osc.output_on(test_ch) + self._osc.set_coupling(test_ch, setup.osc_settings.coupling.value) + self._osc.set_impedance(test_ch, setup.osc_settings.impedance.value) + self._osc.set_vertical(test_ch, setup.osc_settings.full_scale_v, setup.osc_settings.offset_v) + + if run_mode.correction_mode == CorrectionMode.DUAL and setup.channels.osc_ref_ch: + ref_ch = setup.channels.osc_ref_ch + self._osc.output_on(ref_ch) + self._osc.set_coupling(ref_ch, setup.osc_settings.coupling.value) + self._osc.set_impedance(ref_ch, setup.osc_settings.impedance.value) + self._osc.set_vertical(ref_ch, setup.osc_settings.full_scale_v, setup.osc_settings.offset_v) + + if run_mode.trigger_mode == TriggerMode.TRIGGERED: + trig_ch = int(setup.channels.osc_trig_ch or test_ch) + self._osc.output_on(trig_ch) + self._osc.arm_trigger(trig_ch, level_v=0.0) + else: + self._osc.set_free_run() diff --git a/src/app/application/services/sweep/models.py b/src/app/application/services/sweep/models.py new file mode 100644 index 0000000..5d4fe44 --- /dev/null +++ b/src/app/application/services/sweep/models.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +from dataclasses import dataclass, field + +import numpy as np + + +@dataclass(slots=True) +class SweepPlan: + freq_points: np.ndarray + + @property + def total_points(self) -> int: + return int(len(self.freq_points)) + + +@dataclass(slots=True) +class SweepServiceWarning: + code: str + message: str + + +@dataclass(slots=True) +class AcquiredPointData: + actual_freq_hz: float + read_amp_vpp: float + test_times: np.ndarray + test_volts: np.ndarray + ref_times: np.ndarray | None = None + ref_volts: np.ndarray | None = None + warnings: list[SweepServiceWarning] = field(default_factory=list) diff --git a/src/app/application/services/sweep/planner.py b/src/app/application/services/sweep/planner.py new file mode 100644 index 0000000..95f3c98 --- /dev/null +++ b/src/app/application/services/sweep/planner.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +from app.application.services.sweep.models import SweepPlan +from app.domain.models import AppSettings +from app.domain.sweep_engine import compute_sampling_window_s, generate_frequency_points + + +class SweepPlanner: + def plan(self, settings: AppSettings) -> SweepPlan: + return SweepPlan(freq_points=generate_frequency_points(settings.sweep)) + + def compute_sampling_window_s(self, *, freq_hz: float, sample_rate_hz: float, points: int) -> float: + return compute_sampling_window_s(freq_hz=freq_hz, sample_rate_hz=sample_rate_hz, points=points) diff --git a/src/app/application/services/sweep/point_measurement_service.py b/src/app/application/services/sweep/point_measurement_service.py new file mode 100644 index 0000000..3bec594 --- /dev/null +++ b/src/app/application/services/sweep/point_measurement_service.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +from app.application.services.sweep.models import AcquiredPointData +from app.domain.enums import CorrectionMode, TriggerMode +from app.domain.models import AppSettings, SweepPoint +from app.domain.signal_processing import calc_vin_peak, measure_dual_channel, measure_single_channel + + +class PointMeasurementService: + def measure(self, *, settings: AppSettings, acquired: AcquiredPointData) -> SweepPoint: + setup = settings.setup + run_mode = settings.run_mode + + if run_mode.correction_mode == CorrectionMode.DUAL: + gain_linear, gain_db, phase_deg, gain_complex = measure_dual_channel( + acquired.test_times, + acquired.test_volts, + acquired.ref_times if acquired.ref_times is not None else acquired.test_times, + acquired.ref_volts if acquired.ref_volts is not None else acquired.test_volts, + acquired.actual_freq_hz, + ) + else: + vin_peak = calc_vin_peak( + vpp_panel=acquired.read_amp_vpp, + awg_impedance=setup.awg_settings.impedance.value, + osc_impedance=setup.osc_settings.impedance.value, + ) + gain_linear, gain_db, phase_deg, gain_complex = measure_single_channel( + acquired.test_times, + acquired.test_volts, + acquired.actual_freq_hz, + vin_peak, + compute_phase=run_mode.trigger_mode == TriggerMode.TRIGGERED, + ) + + return SweepPoint( + freq_hz=float(acquired.actual_freq_hz), + gain_linear=float(gain_linear), + gain_db=float(gain_db), + phase_deg=float(phase_deg) if phase_deg is not None else None, + gain_complex=complex(gain_complex) if gain_complex is not None else None, + ) diff --git a/src/app/application/services/sweep/waveform_acquirer.py b/src/app/application/services/sweep/waveform_acquirer.py new file mode 100644 index 0000000..f6b06d6 --- /dev/null +++ b/src/app/application/services/sweep/waveform_acquirer.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import numpy as np + +from app.application.ports.instruments import AwgPort, OscPort +from app.application.services.sweep.models import AcquiredPointData, SweepServiceWarning +from app.application.services.sweep.planner import SweepPlanner +from app.domain.auto_range import AutoRangePolicy +from app.domain.enums import CorrectionMode, TriggerMode +from app.domain.models import AppSettings + + +class WaveformAcquirer: + def __init__( + self, + awg: AwgPort, + osc: OscPort, + planner: SweepPlanner, + auto_range_policy: AutoRangePolicy, + ) -> None: + self._awg = awg + self._osc = osc + self._planner = planner + self._auto_range_policy = auto_range_policy + + def acquire(self, *, target_freq_hz: float, settings: AppSettings) -> AcquiredPointData: + setup = settings.setup + run_mode = settings.run_mode + + warnings: list[SweepServiceWarning] = [] + awg_ch = setup.channels.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): + warnings.append( + SweepServiceWarning( + code="FREQ_MISMATCH", + message=f"Requested {target_freq_hz:.6f} Hz, actual {actual_freq:.6f} Hz", + ) + ) + + requested_amp = float(setup.awg_settings.amplitude_vpp) + read_amp = self._awg.get_amplitude_vpp(awg_ch) + if not np.isclose(read_amp, requested_amp, atol=1e-2, rtol=1e-3): + warnings.append( + SweepServiceWarning( + code="AMP_MISMATCH", + message=f"Requested {requested_amp:.6f} Vpp, actual {read_amp:.6f} Vpp", + ) + ) + + sample_rate = self._osc.get_sample_rate() + window_s = self._planner.compute_sampling_window_s( + freq_hz=actual_freq, + sample_rate_hz=sample_rate, + points=setup.osc_settings.points, + ) + + self._osc.set_timebase(window_s) + triggered = run_mode.trigger_mode == TriggerMode.TRIGGERED + self._osc.single_acquire(triggered=triggered) + + test_ch = setup.channels.osc_test_ch + times_t, volts_t = self._osc.read_waveform(test_ch, setup.osc_settings.points) + + if run_mode.auto_range: + current_range, current_offset = self._osc.get_vertical(test_ch) + decision = self._auto_range_policy.decide( + volts=volts_t, + current_range_v=current_range, + current_offset_v=current_offset, + requested_offset_v=setup.osc_settings.offset_v, + ) + if decision.changed: + self._osc.set_vertical(test_ch, decision.target_range_v, decision.target_offset_v) + self._osc.single_acquire(triggered=triggered) + times_t, volts_t = self._osc.read_waveform(test_ch, setup.osc_settings.points) + + ref_times = None + ref_volts = None + if run_mode.correction_mode == CorrectionMode.DUAL: + ref_ch = int(setup.channels.osc_ref_ch or test_ch) + ref_times, ref_volts = self._osc.read_waveform(ref_ch, setup.osc_settings.points) + + return AcquiredPointData( + actual_freq_hz=float(actual_freq), + read_amp_vpp=float(read_amp), + test_times=np.asarray(times_t, dtype=float), + test_volts=np.asarray(volts_t, dtype=float), + ref_times=None if ref_times is None else np.asarray(ref_times, dtype=float), + ref_volts=None if ref_volts is None else np.asarray(ref_volts, dtype=float), + warnings=warnings, + ) diff --git a/src/app/application/services/sweep_task_runner.py b/src/app/application/services/sweep_task_runner.py new file mode 100644 index 0000000..818a97b --- /dev/null +++ b/src/app/application/services/sweep_task_runner.py @@ -0,0 +1,130 @@ +from __future__ import annotations + +import threading +from collections.abc import Callable +from pathlib import Path + +from app.application.dto import SaveTarget, StartSweepCommand +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 +from app.application.use_cases.start_sweep import StartSweepUseCase +from app.application.use_cases.stop_sweep import StopSweepUseCase +from app.domain.models import AppSettings + + +class SweepTaskRunner: + def __init__( + self, + *, + emitter: EventEmitter, + save_measurement_use_case: SaveMeasurementUseCase, + auto_save_dir: Path, + ports_factory: InstrumentPortsFactory, + use_case_factory: Callable[..., StartSweepUseCase] = StartSweepUseCase, + ) -> None: + self._emitter = emitter + self._save_measurement_use_case = save_measurement_use_case + self._auto_save_dir = auto_save_dir + self._ports_factory = ports_factory + self._use_case_factory = use_case_factory + + self._ports: InstrumentPorts | None = None + self._ports_lock = threading.Lock() + self._sweep_thread: threading.Thread | None = None + self._stop_use_case: StopSweepUseCase | None = None + self._active_awg_channel: int | None = None + + def is_running(self) -> bool: + return self._sweep_thread is not None and self._sweep_thread.is_alive() + + def start( + self, + *, + settings: AppSettings, + calibration_enabled: bool, + reference_interpolator: object | None, + ) -> None: + 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() + + def stop(self) -> None: + if self._stop_use_case is not None: + self._stop_use_case.stop() + + def wait(self, timeout: float | None = None) -> None: + if self._sweep_thread is not None: + self._sweep_thread.join(timeout=timeout) + + def shutdown(self, timeout: float = 2.0) -> None: + self.stop() + self.wait(timeout=timeout) + 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.", + ) + 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))) + finally: + self._close_ports() + + def _close_ports(self) -> None: + with self._ports_lock: + ports = self._ports + awg_channel = self._active_awg_channel + self._ports = None + self._active_awg_channel = None + + if ports is None: + return + + 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)) + + try: + ports.awg.close() + except Exception as exc: # noqa: BLE001 + self._emit_warning(code="AWG_CLOSE_FAILED", message=str(exc)) + + try: + ports.osc.close() + except Exception as exc: # noqa: BLE001 + self._emit_warning(code="OSC_CLOSE_FAILED", message=str(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/load_measurement.py b/src/app/application/use_cases/load_measurement.py index 766f35e..30552b1 100644 --- a/src/app/application/use_cases/load_measurement.py +++ b/src/app/application/use_cases/load_measurement.py @@ -1,7 +1,7 @@ from __future__ import annotations from app.application.dto import LoadedMeasurement -from app.infrastructure.persistence.repository_ports import MeasurementRepository +from app.application.ports.persistence import MeasurementRepository class LoadMeasurementUseCase: diff --git a/src/app/application/use_cases/load_reference.py b/src/app/application/use_cases/load_reference.py index 6689ac0..f46c695 100644 --- a/src/app/application/use_cases/load_reference.py +++ b/src/app/application/use_cases/load_reference.py @@ -1,8 +1,8 @@ from __future__ import annotations +from app.application.ports.persistence import ReferenceRepository from app.domain.calibration import build_reference_interpolator from app.domain.models import ReferenceCurve -from app.infrastructure.persistence.repository_ports import ReferenceRepository class LoadReferenceUseCase: diff --git a/src/app/application/use_cases/save_measurement.py b/src/app/application/use_cases/save_measurement.py index 557afb7..49fe1c4 100644 --- a/src/app/application/use_cases/save_measurement.py +++ b/src/app/application/use_cases/save_measurement.py @@ -1,8 +1,8 @@ from __future__ import annotations from app.application.dto import SaveArtifacts, SaveTarget +from app.application.ports.persistence import MeasurementRepository from app.domain.models import AppSettings, SweepResult -from app.infrastructure.persistence.repository_ports import MeasurementRepository class SaveMeasurementUseCase: diff --git a/src/app/application/use_cases/settings_use_case.py b/src/app/application/use_cases/settings_use_case.py index b9c98c8..d11f25d 100644 --- a/src/app/application/use_cases/settings_use_case.py +++ b/src/app/application/use_cases/settings_use_case.py @@ -1,7 +1,7 @@ from __future__ import annotations +from app.application.ports.persistence import SettingsRepository from app.domain.models import AppSettings -from app.infrastructure.persistence.repository_ports import SettingsRepository class SettingsUseCase: diff --git a/src/app/application/use_cases/start_sweep.py b/src/app/application/use_cases/start_sweep.py index 08415a5..b0fb1e8 100644 --- a/src/app/application/use_cases/start_sweep.py +++ b/src/app/application/use_cases/start_sweep.py @@ -4,33 +4,56 @@ import time from datetime import datetime, timezone -import numpy as np - from app.application.dto import StartSweepCommand from app.application.events import ( EventEmitter, SweepCompleted, - SweepDataUpdated, SweepFailed, + SweepDataUpdated, SweepProgress, SweepStarted, SweepStopped, SweepWarning, ) -from app.domain.calibration import apply_reference_to_point -from app.domain.enums import CorrectionMode, TriggerMode -from app.domain.models import SweepPoint, SweepResult -from app.domain.signal_processing import calc_vin_peak, measure_dual_channel, measure_single_channel -from app.domain.sweep_engine import compute_sampling_window_s, generate_frequency_points +from app.application.ports.instruments import AwgPort, OscPort +from app.application.services.sweep import ( + CalibrationApplier, + InstrumentConfigurator, + PointMeasurementService, + SweepPlanner, + WaveformAcquirer, +) +from app.domain.auto_range import AutoRangePolicy +from app.domain.models import SweepResult from app.domain.validators import ValidationError, validate_settings -from app.infrastructure.instruments.ports import AwgPort, OscPort class StartSweepUseCase: - def __init__(self, awg: AwgPort, osc: OscPort, stop_event: threading.Event) -> None: + def __init__( + self, + awg: AwgPort, + osc: OscPort, + stop_event: threading.Event, + *, + planner: SweepPlanner | None = None, + configurator: InstrumentConfigurator | None = None, + acquirer: WaveformAcquirer | None = None, + measurement_service: PointMeasurementService | None = None, + calibration_applier: CalibrationApplier | None = None, + ) -> None: self._awg = awg self._osc = osc self._stop_event = stop_event + self._planner = planner or SweepPlanner() + self._configurator = configurator or InstrumentConfigurator(awg=awg, osc=osc) + self._acquirer = acquirer or WaveformAcquirer( + awg=awg, + osc=osc, + planner=self._planner, + auto_range_policy=AutoRangePolicy(), + ) + self._measurement_service = measurement_service or PointMeasurementService() + self._calibration_applier = calibration_applier or CalibrationApplier() def run(self, cmd: StartSweepCommand, emitter: EventEmitter) -> SweepResult: try: @@ -44,111 +67,30 @@ def run(self, cmd: StartSweepCommand, emitter: EventEmitter) -> SweepResult: ) settings = cmd.settings - sweep = settings.sweep - setup = settings.setup - run_mode = settings.run_mode + plan = self._planner.plan(settings) + emitter.emit(SweepStarted(total_points=plan.total_points)) - freq_points = generate_frequency_points(sweep) - emitter.emit(SweepStarted(total_points=len(freq_points))) + self._configurator.configure(settings) + emitter.emit(SweepWarning(code="READY", message="Instruments configured")) - self._configure_instruments(cmd, emitter) - - for index, target_freq in enumerate(freq_points, start=1): + for index, target_freq in enumerate(plan.freq_points, start=1): if self._stop_event.is_set(): result.meta["stopped_at"] = datetime.now(timezone.utc).isoformat(timespec="seconds") emitter.emit(SweepStopped(result=result)) return result - awg_ch = setup.channels.awg_ch - self._awg.set_frequency(float(target_freq), awg_ch) - actual_freq = self._awg.get_frequency(awg_ch) - - if not np.isclose(actual_freq, target_freq, atol=1e-3, rtol=5e-6): - emitter.emit( - SweepWarning( - code="FREQ_MISMATCH", - message=( - f"Requested {target_freq:.6f} Hz, actual {actual_freq:.6f} Hz" - ), - ) - ) - - requested_amp = float(setup.awg_settings.amplitude_vpp) - read_amp = self._awg.get_amplitude_vpp(awg_ch) - if not np.isclose(read_amp, requested_amp, atol=1e-2, rtol=1e-3): - emitter.emit( - SweepWarning( - code="AMP_MISMATCH", - message=( - f"Requested {requested_amp:.6f} Vpp, actual {read_amp:.6f} Vpp" - ), - ) - ) - - sample_rate = self._osc.get_sample_rate() - window_s = compute_sampling_window_s( - freq_hz=actual_freq, - sample_rate_hz=sample_rate, - points=setup.osc_settings.points, - ) + acquired = self._acquirer.acquire(target_freq_hz=float(target_freq), settings=settings) + for warning in acquired.warnings: + emitter.emit(SweepWarning(code=warning.code, message=warning.message)) - self._osc.set_timebase(window_s) - triggered = run_mode.trigger_mode == TriggerMode.TRIGGERED - self._osc.single_acquire(triggered=triggered) - - test_ch = setup.channels.osc_test_ch - times_t, volts_t = self._osc.read_waveform(test_ch, setup.osc_settings.points) - - if run_mode.auto_range and self._adjust_auto_range(test_ch, volts_t, setup.osc_settings.offset_v): - self._osc.single_acquire(triggered=triggered) - times_t, volts_t = self._osc.read_waveform(test_ch, setup.osc_settings.points) - - if run_mode.correction_mode == CorrectionMode.DUAL: - ref_ch = int(setup.channels.osc_ref_ch or test_ch) - times_r, volts_r = self._osc.read_waveform(ref_ch, setup.osc_settings.points) - gain_linear, gain_db, phase_deg, gain_complex = measure_dual_channel( - times_t, - volts_t, - times_r, - volts_r, - actual_freq, - ) - else: - vin_peak = calc_vin_peak( - vpp_panel=read_amp, - awg_impedance=setup.awg_settings.impedance.value, - osc_impedance=setup.osc_settings.impedance.value, - ) - gain_linear, gain_db, phase_deg, gain_complex = measure_single_channel( - times_t, - volts_t, - actual_freq, - vin_peak, - compute_phase=triggered, - ) - - point = SweepPoint( - freq_hz=float(actual_freq), - gain_linear=float(gain_linear), - gain_db=float(gain_db), - phase_deg=float(phase_deg) if phase_deg is not None else None, - gain_complex=complex(gain_complex) if gain_complex is not None else None, - ) - - if cmd.calibration_enabled and cmd.reference_interpolator is not None: - use_phase = ( - run_mode.correction_mode == CorrectionMode.DUAL - or run_mode.trigger_mode == TriggerMode.TRIGGERED - ) - ref_value = cmd.reference_interpolator(np.array([point.freq_hz]))[0] - point = apply_reference_to_point(point, ref_value, use_phase=use_phase) + point = self._measurement_service.measure(settings=settings, acquired=acquired) + point = self._calibration_applier.apply(point=point, cmd=cmd) result.append(point) - emitter.emit(SweepProgress(freq_hz=point.freq_hz, point_index=index, total_points=len(freq_points))) + emitter.emit(SweepProgress(freq_hz=point.freq_hz, point_index=index, total_points=plan.total_points)) emitter.emit(SweepDataUpdated(last_point=point, partial_result=result)) - # Give stop signals a chance to be observed in long hardware loops. time.sleep(0.001) result.meta["completed_at"] = datetime.now(timezone.utc).isoformat(timespec="seconds") @@ -161,77 +103,3 @@ def run(self, cmd: StartSweepCommand, emitter: EventEmitter) -> SweepResult: except Exception as exc: # noqa: BLE001 emitter.emit(SweepFailed(error_code="SWEEP_RUNTIME", message=str(exc))) return SweepResult() - - def _configure_instruments(self, cmd: StartSweepCommand, emitter: EventEmitter) -> None: - settings = cmd.settings - setup = settings.setup - run_mode = settings.run_mode - - awg_ch = setup.channels.awg_ch - test_ch = setup.channels.osc_test_ch - - if run_mode.auto_reset: - self._awg.reset() - self._osc.reset() - - self._awg.output_on(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) - - self._osc.output_on(test_ch) - self._osc.set_coupling(test_ch, setup.osc_settings.coupling.value) - self._osc.set_impedance(test_ch, setup.osc_settings.impedance.value) - self._osc.set_vertical(test_ch, setup.osc_settings.full_scale_v, setup.osc_settings.offset_v) - - if run_mode.correction_mode == CorrectionMode.DUAL and setup.channels.osc_ref_ch: - self._osc.output_on(setup.channels.osc_ref_ch) - self._osc.set_coupling(setup.channels.osc_ref_ch, setup.osc_settings.coupling.value) - self._osc.set_impedance(setup.channels.osc_ref_ch, setup.osc_settings.impedance.value) - self._osc.set_vertical( - setup.channels.osc_ref_ch, - setup.osc_settings.full_scale_v, - setup.osc_settings.offset_v, - ) - - if run_mode.trigger_mode == TriggerMode.TRIGGERED: - trig_ch = int(setup.channels.osc_trig_ch or test_ch) - self._osc.output_on(trig_ch) - self._osc.arm_trigger(trig_ch, level_v=0.0) - else: - self._osc.set_free_run() - - emitter.emit(SweepWarning(code="READY", message="Instruments configured")) - - def _adjust_auto_range(self, channel: int, volts: np.ndarray, requested_offset_v: float) -> bool: - if volts is None or len(volts) == 0: - return False - - vmax = float(np.max(volts)) - vmin = float(np.min(volts)) - vpp = vmax - vmin - midpoint = (vmax + vmin) / 2.0 - - current_range, current_offset = self._osc.get_vertical(channel) - if current_range <= 0: - return False - - ratio = vpp / current_range - target_range = current_range - target_offset = requested_offset_v - - if ratio > 0.85: - target_range = vpp / 0.7 - elif 0.0 < ratio < 0.55: - target_range = max(vpp / 0.7, current_range * 0.5) - - if abs(midpoint - current_offset) > (current_range * 0.2): - target_offset = midpoint - - range_changed = not np.isclose(target_range, current_range, rtol=1e-2, atol=1e-3) - offset_changed = not np.isclose(target_offset, current_offset, rtol=1e-2, atol=1e-3) - - if range_changed or offset_changed: - self._osc.set_vertical(channel, float(target_range), float(target_offset)) - return True - - return False diff --git a/src/app/bootstrap.py b/src/app/bootstrap.py new file mode 100644 index 0000000..c5f3e26 --- /dev/null +++ b/src/app/bootstrap.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from app.application.use_cases.load_measurement import LoadMeasurementUseCase +from app.application.use_cases.load_reference import LoadReferenceUseCase +from app.application.use_cases.save_measurement import SaveMeasurementUseCase +from app.application.use_cases.settings_use_case import SettingsUseCase +from app.infrastructure.instruments.equips_factory import create_instrument_ports, resolve_visa_address +from app.infrastructure.instruments.resource_scanner import PyVisaResourceScanner +from app.infrastructure.persistence.measurement_repo_mat_csv import MatCsvMeasurementRepository +from app.infrastructure.persistence.reference_repo_mat import MatReferenceRepository +from app.infrastructure.persistence.settings_repo_json import JsonSettingsRepository +from app.presentation.tk.app_window import AppWindow +from app.presentation.tk.controller import TkController +from app.runtime.paths import AppPaths + + +@dataclass(slots=True) +class DesktopApp: + window: AppWindow + controller: TkController + paths: AppPaths + + def run(self) -> None: + self.controller.initialize() + self.window.mainloop() + + +def build_desktop_app(paths: AppPaths | None = None) -> DesktopApp: + app_paths = paths or AppPaths.default() + window = AppWindow() + vm = window.vm + + settings_repo = JsonSettingsRepository(config_path=app_paths.settings_path) + measurement_repo = MatCsvMeasurementRepository() + reference_repo = MatReferenceRepository() + save_measurement_use_case = SaveMeasurementUseCase(measurement_repo) + + controller = TkController( + window=window, + vm=vm, + settings_use_case=SettingsUseCase(settings_repo), + save_measurement_use_case=save_measurement_use_case, + load_measurement_use_case=LoadMeasurementUseCase(measurement_repo), + load_reference_use_case=LoadReferenceUseCase(reference_repo), + scanner=PyVisaResourceScanner(), + ports_factory=create_instrument_ports, + resolve_address=resolve_visa_address, + paths=app_paths, + ) + return DesktopApp(window=window, controller=controller, paths=app_paths) + + +def run_desktop_app(paths: AppPaths | None = None) -> None: + build_desktop_app(paths=paths).run() + diff --git a/src/app/domain/auto_range.py b/src/app/domain/auto_range.py new file mode 100644 index 0000000..079a3df --- /dev/null +++ b/src/app/domain/auto_range.py @@ -0,0 +1,67 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np + + +@dataclass(slots=True) +class AutoRangeDecision: + changed: bool + target_range_v: float + target_offset_v: float + + +class AutoRangePolicy: + def __init__( + self, + *, + high_threshold: float = 0.85, + low_threshold: float = 0.55, + target_fill_ratio: float = 0.7, + offset_threshold_ratio: float = 0.2, + ) -> None: + self._high_threshold = high_threshold + self._low_threshold = low_threshold + self._target_fill_ratio = target_fill_ratio + self._offset_threshold_ratio = offset_threshold_ratio + + def decide( + self, + *, + volts: np.ndarray, + current_range_v: float, + current_offset_v: float, + requested_offset_v: float, + ) -> AutoRangeDecision: + if volts is None or len(volts) == 0 or current_range_v <= 0: + return AutoRangeDecision( + changed=False, + target_range_v=float(current_range_v), + target_offset_v=float(current_offset_v), + ) + + vmax = float(np.max(volts)) + vmin = float(np.min(volts)) + vpp = vmax - vmin + midpoint = (vmax + vmin) / 2.0 + + target_range = float(current_range_v) + target_offset = float(requested_offset_v) + ratio = vpp / current_range_v + + if ratio > self._high_threshold: + target_range = vpp / self._target_fill_ratio + elif 0.0 < ratio < self._low_threshold: + target_range = max(vpp / self._target_fill_ratio, current_range_v * 0.5) + + if abs(midpoint - current_offset_v) > (current_range_v * self._offset_threshold_ratio): + target_offset = midpoint + + range_changed = not np.isclose(target_range, current_range_v, rtol=1e-2, atol=1e-3) + offset_changed = not np.isclose(target_offset, current_offset_v, rtol=1e-2, atol=1e-3) + return AutoRangeDecision( + changed=bool(range_changed or offset_changed), + target_range_v=float(target_range), + target_offset_v=float(target_offset), + ) diff --git a/src/app/domain/calibration.py b/src/app/domain/calibration.py index 22729da..5445f26 100644 --- a/src/app/domain/calibration.py +++ b/src/app/domain/calibration.py @@ -91,7 +91,10 @@ def apply_reference_to_point(point: SweepPoint, ref_value: complex | float, use_ else: raw = complex(point.gain_linear, 0.0) - corrected = raw / complex(ref_value) + ref_complex = complex(ref_value) + if abs(ref_complex) < eps: + ref_complex = complex(eps, 0.0) + corrected = raw / ref_complex gain_linear = float(np.abs(corrected)) gain_db = float(20.0 * np.log10(max(gain_linear, eps))) phase_deg = float(np.degrees(np.angle(corrected))) diff --git a/src/app/domain/exporters.py b/src/app/domain/exporters.py index 614f754..5fe7428 100644 --- a/src/app/domain/exporters.py +++ b/src/app/domain/exporters.py @@ -12,8 +12,18 @@ def result_to_arrays(result: SweepResult) -> dict[str, np.ndarray]: gain = np.array([p.gain_linear for p in result.points], dtype=float) gain_db = np.array([p.gain_db for p in result.points], dtype=float) - phase_values = [p.phase_deg for p in result.points if p.phase_deg is not None] - complex_values = [p.gain_complex for p in result.points if p.gain_complex is not None] + phase_values = np.array( + [np.nan if p.phase_deg is None else float(p.phase_deg) for p in result.points], + dtype=float, + ) + complex_real = np.array( + [np.nan if p.gain_complex is None else float(p.gain_complex.real) for p in result.points], + dtype=float, + ) + complex_imag = np.array( + [np.nan if p.gain_complex is None else float(p.gain_complex.imag) for p in result.points], + dtype=float, + ) arrays: dict[str, np.ndarray] = { "freq_hz": freq, @@ -21,11 +31,11 @@ def result_to_arrays(result: SweepResult) -> dict[str, np.ndarray]: "gain_db": gain_db, } - if phase_values: - arrays["phase_deg"] = np.array(phase_values, dtype=float) - if complex_values: - arrays["gain_complex_real"] = np.array([v.real for v in complex_values], dtype=float) - arrays["gain_complex_imag"] = np.array([v.imag for v in complex_values], dtype=float) + if np.any(~np.isnan(phase_values)): + arrays["phase_deg"] = phase_values + if np.any(~np.isnan(complex_real)) or np.any(~np.isnan(complex_imag)): + arrays["gain_complex_real"] = complex_real + arrays["gain_complex_imag"] = complex_imag return arrays diff --git a/src/app/infrastructure/instruments/awg_adapter.py b/src/app/infrastructure/instruments/awg_adapter.py index 4c256da..15d165b 100644 --- a/src/app/infrastructure/instruments/awg_adapter.py +++ b/src/app/infrastructure/instruments/awg_adapter.py @@ -1,23 +1,20 @@ from __future__ import annotations +from app.infrastructure.instruments.vendor_gateway import create_vendor_instrument + + class EquipsAwgAdapter: def __init__(self, model: str, visa_address: str) -> None: - try: - from equips import inst_mapping - except ModuleNotFoundError as exc: - raise RuntimeError( - "Missing runtime dependency: pyvisa/equips is required for instrument access" - ) from exc - - if model not in inst_mapping: - raise ValueError(f"Unsupported AWG model: {model}") - self._inst = inst_mapping[model](name=model, visa_address=visa_address) + self._inst = create_vendor_instrument(model=model, visa_address=visa_address) def reset(self) -> None: self._inst.rst() def output_on(self, channel: int) -> None: - self._inst.set_on(ch=channel) + self._inst.set_on(on=True, ch=channel) + + def output_off(self, channel: int) -> None: + self._inst.set_on(on=False, ch=channel) def set_impedance(self, mode: str, channel: int) -> None: self._inst.set_imp(imp=mode, ch=channel) diff --git a/src/app/infrastructure/instruments/equips_factory.py b/src/app/infrastructure/instruments/equips_factory.py index 25bae2a..d563d77 100644 --- a/src/app/infrastructure/instruments/equips_factory.py +++ b/src/app/infrastructure/instruments/equips_factory.py @@ -1,21 +1,10 @@ from __future__ import annotations -from dataclasses import dataclass - +from app.application.ports.instruments import InstrumentPorts from app.domain.enums import ConnectionMode from app.domain.models import InstrumentEndpoint, InstrumentSetup from app.infrastructure.instruments.awg_adapter import EquipsAwgAdapter from app.infrastructure.instruments.osc_adapter import EquipsOscAdapter -from app.infrastructure.instruments.ports import AwgPort, OscPort - - -@dataclass(slots=True) -class InstrumentPorts: - awg: AwgPort - osc: OscPort - awg_address: str - osc_address: str - def resolve_visa_address(endpoint: InstrumentEndpoint) -> str: diff --git a/src/app/infrastructure/instruments/osc_adapter.py b/src/app/infrastructure/instruments/osc_adapter.py index 51827fd..2422bb6 100644 --- a/src/app/infrastructure/instruments/osc_adapter.py +++ b/src/app/infrastructure/instruments/osc_adapter.py @@ -2,18 +2,12 @@ import numpy as np +from app.infrastructure.instruments.vendor_gateway import create_vendor_instrument + + class EquipsOscAdapter: def __init__(self, model: str, visa_address: str) -> None: - try: - from equips import inst_mapping - except ModuleNotFoundError as exc: - raise RuntimeError( - "Missing runtime dependency: pyvisa/equips is required for instrument access" - ) from exc - - if model not in inst_mapping: - raise ValueError(f"Unsupported OSC model: {model}") - self._inst = inst_mapping[model](name=model, visa_address=visa_address) + self._inst = create_vendor_instrument(model=model, visa_address=visa_address) def reset(self) -> None: self._inst.rst() diff --git a/src/app/infrastructure/instruments/ports.py b/src/app/infrastructure/instruments/ports.py index cdc027b..c5fb368 100644 --- a/src/app/infrastructure/instruments/ports.py +++ b/src/app/infrastructure/instruments/ports.py @@ -1,36 +1,3 @@ -from __future__ import annotations +from app.application.ports.instruments import AwgPort, InstrumentPorts, InstrumentPortsFactory, OscPort, ResourceScannerPort -from typing import Protocol - -import numpy as np - - -class AwgPort(Protocol): - def reset(self) -> None: ... - def output_on(self, channel: int) -> None: ... - def set_impedance(self, mode: str, channel: int) -> None: ... - def set_frequency(self, hz: float, channel: int) -> None: ... - def get_frequency(self, channel: int) -> float: ... - def set_amplitude_vpp(self, vpp: float, channel: int) -> None: ... - def get_amplitude_vpp(self, channel: int) -> float: ... - def close(self) -> None: ... - - -class OscPort(Protocol): - def reset(self) -> None: ... - def output_on(self, channel: int) -> None: ... - def set_timebase(self, window_s: float, offset_s: float | None = None) -> None: ... - def set_vertical(self, channel: int, full_scale_v: float, offset_v: float) -> None: ... - def get_vertical(self, channel: int) -> tuple[float, float]: ... - def set_coupling(self, channel: int, mode: str) -> None: ... - def set_impedance(self, channel: int, mode: str) -> None: ... - def arm_trigger(self, channel: int, level_v: float) -> None: ... - def set_free_run(self) -> None: ... - def single_acquire(self, triggered: bool) -> None: ... - def read_waveform(self, channel: int, points: int | None) -> tuple[np.ndarray, np.ndarray]: ... - def get_sample_rate(self) -> float: ... - def close(self) -> None: ... - - -class ResourceScannerPort(Protocol): - def list_resources(self) -> tuple[str, ...]: ... +__all__ = ["AwgPort", "OscPort", "ResourceScannerPort", "InstrumentPorts", "InstrumentPortsFactory"] diff --git a/src/app/infrastructure/instruments/vendor_gateway.py b/src/app/infrastructure/instruments/vendor_gateway.py new file mode 100644 index 0000000..b82d3c6 --- /dev/null +++ b/src/app/infrastructure/instruments/vendor_gateway.py @@ -0,0 +1,12 @@ +from __future__ import annotations + + +def create_vendor_instrument(model: str, visa_address: str) -> object: + try: + from equips import inst_mapping + except ModuleNotFoundError as exc: + raise RuntimeError("Missing runtime dependency: pyvisa/equips is required for instrument access") from exc + + if model not in inst_mapping: + raise ValueError(f"Unsupported instrument model: {model}") + return inst_mapping[model](name=model, visa_address=visa_address) diff --git a/src/app/infrastructure/persistence/file_store.py b/src/app/infrastructure/persistence/file_store.py new file mode 100644 index 0000000..bc2eca3 --- /dev/null +++ b/src/app/infrastructure/persistence/file_store.py @@ -0,0 +1,18 @@ +from __future__ import annotations + +from pathlib import Path + + +class FileStore: + def __init__(self, path: Path) -> None: + self.path = path + + def exists(self) -> bool: + return self.path.exists() + + def read_text(self, *, encoding: str = "utf-8") -> str: + return self.path.read_text(encoding=encoding) + + def write_text(self, data: str, *, encoding: str = "utf-8") -> None: + self.path.parent.mkdir(parents=True, exist_ok=True) + self.path.write_text(data, encoding=encoding) diff --git a/src/app/infrastructure/persistence/measurement_exporter.py b/src/app/infrastructure/persistence/measurement_exporter.py new file mode 100644 index 0000000..cc1ba27 --- /dev/null +++ b/src/app/infrastructure/persistence/measurement_exporter.py @@ -0,0 +1,92 @@ +from __future__ import annotations + +import csv +import json +from datetime import datetime +from pathlib import Path + +import numpy as np +from scipy.io import savemat + +from app.application.dto import SaveArtifacts, SaveTarget +from app.domain.exporters import result_to_arrays, settings_to_metadata +from app.domain.models import AppSettings, SweepResult + + +class MeasurementExporter: + def export(self, result: SweepResult, settings: AppSettings, target: SaveTarget) -> SaveArtifacts: + directory = target.base_path.parent + directory.mkdir(parents=True, exist_ok=True) + + stem = target.base_path.stem if target.base_path.suffix else target.base_path.name + prefix = datetime.now().strftime("%Y%m%d_%H_%M_%S_") if target.include_timestamp else "" + file_base = f"{prefix}{stem}" + + mat_path = directory / f"{file_base}.mat" + csv_path = directory / f"{file_base}.csv" + txt_path = directory / f"{file_base}.txt" + + arrays = result_to_arrays(result) + payload: dict[str, object] = { + "schema_version": settings.schema_version, + "metadata_json": json.dumps(settings_to_metadata(settings), ensure_ascii=True), + } + payload.update(arrays) + savemat(mat_path, payload) + + freq = arrays.get("freq_hz", np.array([], dtype=float)) + gain_linear = arrays.get("gain_linear", np.array([], dtype=float)) + gain_db = arrays.get("gain_db", np.array([], dtype=float)) + phase = arrays.get("phase_deg", np.array([], dtype=float)) + + headers = ["freq_hz", "gain_linear", "gain_db", "phase_deg"] + with csv_path.open("w", newline="", encoding="utf-8") as fh: + writer = csv.writer(fh) + writer.writerow(headers) + for idx in range(len(freq)): + writer.writerow( + [ + float(freq[idx]), + float(gain_linear[idx]) if idx < len(gain_linear) else "", + float(gain_db[idx]) if idx < len(gain_db) else "", + _optional_float(phase, idx), + ] + ) + + rows = np.column_stack( + [ + freq, + gain_linear if len(gain_linear) == len(freq) else np.full(len(freq), np.nan), + gain_db if len(gain_db) == len(freq) else np.full(len(freq), np.nan), + phase if len(phase) == len(freq) else np.full(len(freq), np.nan), + ] + ) + np.savetxt(txt_path, rows, delimiter="\t", header="\t".join(headers), comments="") + + gain_plot_path = None + db_plot_path = None + gain_fig = target.figures.get("gain") + db_fig = target.figures.get("db") + if gain_fig is not None: + gain_plot_path = directory / f"{file_base}_gain.png" + gain_fig.savefig(gain_plot_path, dpi=300) + if db_fig is not None: + db_plot_path = directory / f"{file_base}_gain_db.png" + db_fig.savefig(db_plot_path, dpi=300) + + return SaveArtifacts( + mat_path=mat_path, + csv_path=csv_path, + txt_path=txt_path, + gain_plot_path=gain_plot_path, + db_plot_path=db_plot_path, + ) + + +def _optional_float(values: np.ndarray, idx: int) -> float | str: + if idx >= len(values): + return "" + value = float(values[idx]) + if np.isnan(value): + return "" + return value diff --git a/src/app/infrastructure/persistence/measurement_loader.py b/src/app/infrastructure/persistence/measurement_loader.py new file mode 100644 index 0000000..822184f --- /dev/null +++ b/src/app/infrastructure/persistence/measurement_loader.py @@ -0,0 +1,94 @@ +from __future__ import annotations + +import csv +from pathlib import Path + +import numpy as np +from scipy.io import loadmat + +from app.application.dto import LoadedMeasurement +from app.domain.models import SweepPoint, SweepResult + + +class MeasurementLoader: + def load(self, file_path: str) -> LoadedMeasurement: + path = Path(file_path) + suffix = path.suffix.lower() + + if suffix == ".mat": + payload = loadmat(path) + freq = self._get_array(payload, ["freq_hz", "freq"]) # type: ignore[arg-type] + gain_db = self._get_array(payload, ["gain_db", "gain_db_raw", "gain_db_corr"]) + gain_linear = self._get_array(payload, ["gain_linear", "gain_raw"], required=False) + if gain_linear is None: + gain_linear = np.power(10.0, gain_db / 20.0) + phase = self._get_array(payload, ["phase_deg", "phase", "phase_deg_corr", "phase_corr"], required=False) + elif suffix == ".csv": + freq_l: list[float] = [] + gain_l: list[float] = [] + gain_db_l: list[float] = [] + phase_l: list[float] = [] + with path.open("r", encoding="utf-8") as fh: + reader = csv.DictReader(fh) + for row in reader: + freq_l.append(float(row.get("freq_hz", "0") or 0.0)) + gain_l.append(float(row.get("gain_linear", "0") or 0.0)) + gain_db_l.append(float(row.get("gain_db", "0") or 0.0)) + phase_value = row.get("phase_deg") + phase_l.append(float(phase_value) if phase_value not in (None, "") else np.nan) + freq = np.array(freq_l, dtype=float) + gain_linear = np.array(gain_l, dtype=float) + gain_db = np.array(gain_db_l, dtype=float) + phase_values = np.array(phase_l, dtype=float) if phase_l else None + phase = phase_values if phase_values is not None and np.any(~np.isnan(phase_values)) else None + payload = {"freq_hz": freq, "gain_linear": gain_linear, "gain_db": gain_db} + if phase is not None: + payload["phase_deg"] = phase + else: + raise ValueError(f"Unsupported file type: {suffix}") + + points: list[SweepPoint] = [] + for idx in range(len(freq)): + phase_deg = _optional_phase(phase, idx) + points.append( + SweepPoint( + freq_hz=float(freq[idx]), + gain_linear=float(gain_linear[idx]) if idx < len(gain_linear) else 0.0, + gain_db=float(gain_db[idx]) if idx < len(gain_db) else 0.0, + phase_deg=phase_deg, + gain_complex=( + complex( + float(gain_linear[idx]) * np.cos(np.deg2rad(phase_deg)), + float(gain_linear[idx]) * np.sin(np.deg2rad(phase_deg)), + ) + if phase_deg is not None and idx < len(gain_linear) + else None + ), + ) + ) + + return LoadedMeasurement(result=SweepResult(points=points), raw_payload=payload) + + def _get_array( + self, + payload: dict[str, np.ndarray], + keys: list[str], + *, + required: bool = True, + ) -> np.ndarray | None: + for key in keys: + value = payload.get(key) + if isinstance(value, np.ndarray): + return np.atleast_1d(np.asarray(value, dtype=float).squeeze()) + if required: + raise ValueError(f"Missing required keys: {keys}") + return None + + +def _optional_phase(phase: np.ndarray | None, idx: int) -> float | None: + if phase is None or idx >= len(phase): + return None + value = float(phase[idx]) + if np.isnan(value): + return None + return value diff --git a/src/app/infrastructure/persistence/measurement_repo_mat_csv.py b/src/app/infrastructure/persistence/measurement_repo_mat_csv.py index 057285a..edf6f98 100644 --- a/src/app/infrastructure/persistence/measurement_repo_mat_csv.py +++ b/src/app/infrastructure/persistence/measurement_repo_mat_csv.py @@ -1,158 +1,23 @@ from __future__ import annotations -import csv -import json -from datetime import datetime -from pathlib import Path - -import numpy as np -from scipy.io import loadmat, savemat - from app.application.dto import LoadedMeasurement, SaveArtifacts, SaveTarget -from app.domain.exporters import result_to_arrays, settings_to_metadata -from app.domain.models import AppSettings, SweepPoint, SweepResult +from app.domain.models import AppSettings, SweepResult +from app.infrastructure.persistence.measurement_exporter import MeasurementExporter +from app.infrastructure.persistence.measurement_loader import MeasurementLoader class MatCsvMeasurementRepository: - def save(self, result: SweepResult, settings: AppSettings, target: SaveTarget) -> SaveArtifacts: - directory = target.base_path.parent - directory.mkdir(parents=True, exist_ok=True) - - stem = target.base_path.stem if target.base_path.suffix else target.base_path.name - prefix = datetime.now().strftime("%Y%m%d_%H_%M_%S_") if target.include_timestamp else "" - file_base = f"{prefix}{stem}" - - mat_path = directory / f"{file_base}.mat" - csv_path = directory / f"{file_base}.csv" - txt_path = directory / f"{file_base}.txt" - - arrays = result_to_arrays(result) - payload: dict[str, object] = { - "schema_version": settings.schema_version, - "metadata_json": json.dumps(settings_to_metadata(settings), ensure_ascii=True), - } - payload.update(arrays) - savemat(mat_path, payload) - - freq = arrays.get("freq_hz", np.array([], dtype=float)) - gain_linear = arrays.get("gain_linear", np.array([], dtype=float)) - gain_db = arrays.get("gain_db", np.array([], dtype=float)) - phase = arrays.get("phase_deg", np.array([], dtype=float)) - - headers = ["freq_hz", "gain_linear", "gain_db", "phase_deg"] - with csv_path.open("w", newline="", encoding="utf-8") as fh: - writer = csv.writer(fh) - writer.writerow(headers) - for idx in range(len(freq)): - writer.writerow( - [ - float(freq[idx]), - float(gain_linear[idx]) if idx < len(gain_linear) else "", - float(gain_db[idx]) if idx < len(gain_db) else "", - float(phase[idx]) if idx < len(phase) else "", - ] - ) - - rows = np.column_stack( - [ - freq, - gain_linear if len(gain_linear) == len(freq) else np.full(len(freq), np.nan), - gain_db if len(gain_db) == len(freq) else np.full(len(freq), np.nan), - phase if len(phase) == len(freq) else np.full(len(freq), np.nan), - ] - ) - np.savetxt(txt_path, rows, delimiter="\t", header="\t".join(headers), comments="") - - gain_plot_path = None - db_plot_path = None - - gain_fig = target.figures.get("gain") - db_fig = target.figures.get("db") - - if gain_fig is not None: - gain_plot_path = directory / f"{file_base}_gain.png" - gain_fig.savefig(gain_plot_path, dpi=300) - if db_fig is not None: - db_plot_path = directory / f"{file_base}_gain_db.png" - db_fig.savefig(db_plot_path, dpi=300) + def __init__( + self, + *, + exporter: MeasurementExporter | None = None, + loader: MeasurementLoader | None = None, + ) -> None: + self._exporter = exporter or MeasurementExporter() + self._loader = loader or MeasurementLoader() - return SaveArtifacts( - mat_path=mat_path, - csv_path=csv_path, - txt_path=txt_path, - gain_plot_path=gain_plot_path, - db_plot_path=db_plot_path, - ) + def save(self, result: SweepResult, settings: AppSettings, target: SaveTarget) -> SaveArtifacts: + return self._exporter.export(result=result, settings=settings, target=target) def load(self, file_path: str) -> LoadedMeasurement: - path = Path(file_path) - suffix = path.suffix.lower() - - if suffix == ".mat": - payload = loadmat(path) - freq = self._get_array(payload, ["freq_hz", "freq"]) # type: ignore[arg-type] - gain_linear = self._get_array(payload, ["gain_linear", "gain_raw"]) - gain_db = self._get_array(payload, ["gain_db", "gain_db_raw", "gain_db_corr"]) - phase = self._get_array(payload, ["phase_deg", "phase", "phase_deg_corr"], required=False) - - elif suffix == ".csv": - freq_l: list[float] = [] - gain_l: list[float] = [] - gain_db_l: list[float] = [] - phase_l: list[float] = [] - with path.open("r", encoding="utf-8") as fh: - reader = csv.DictReader(fh) - for row in reader: - freq_l.append(float(row.get("freq_hz", "0") or 0.0)) - gain_l.append(float(row.get("gain_linear", "0") or 0.0)) - gain_db_l.append(float(row.get("gain_db", "0") or 0.0)) - if row.get("phase_deg") not in (None, ""): - phase_l.append(float(row["phase_deg"])) - freq = np.array(freq_l, dtype=float) - gain_linear = np.array(gain_l, dtype=float) - gain_db = np.array(gain_db_l, dtype=float) - phase = np.array(phase_l, dtype=float) if phase_l else None - payload = {"freq_hz": freq, "gain_linear": gain_linear, "gain_db": gain_db} - if phase is not None: - payload["phase_deg"] = phase - - else: - raise ValueError(f"Unsupported file type: {suffix}") - - points: list[SweepPoint] = [] - n = len(freq) - for idx in range(n): - phase_deg = float(phase[idx]) if phase is not None and idx < len(phase) else None - points.append( - SweepPoint( - freq_hz=float(freq[idx]), - gain_linear=float(gain_linear[idx]) if idx < len(gain_linear) else 0.0, - gain_db=float(gain_db[idx]) if idx < len(gain_db) else 0.0, - phase_deg=phase_deg, - gain_complex=( - complex( - float(gain_linear[idx]) * np.cos(np.deg2rad(phase_deg)), - float(gain_linear[idx]) * np.sin(np.deg2rad(phase_deg)), - ) - if phase_deg is not None and idx < len(gain_linear) - else None - ), - ) - ) - - return LoadedMeasurement(result=SweepResult(points=points), raw_payload=payload) - - def _get_array( - self, - payload: dict[str, np.ndarray], - keys: list[str], - *, - required: bool = True, - ) -> np.ndarray | None: - for key in keys: - value = payload.get(key) - if isinstance(value, np.ndarray): - return np.asarray(value, dtype=float).squeeze() - if required: - raise ValueError(f"Missing required keys: {keys}") - return None + return self._loader.load(file_path) diff --git a/src/app/infrastructure/persistence/repository_ports.py b/src/app/infrastructure/persistence/repository_ports.py index 7eaa96f..30d3e0f 100644 --- a/src/app/infrastructure/persistence/repository_ports.py +++ b/src/app/infrastructure/persistence/repository_ports.py @@ -1,20 +1,3 @@ -from __future__ import annotations +from app.application.ports.persistence import MeasurementRepository, ReferenceRepository, SettingsRepository -from typing import Protocol - -from app.application.dto import LoadedMeasurement, SaveArtifacts, SaveTarget -from app.domain.models import AppSettings, ReferenceCurve, SweepResult - - -class SettingsRepository(Protocol): - def load(self) -> AppSettings: ... - def save(self, settings: AppSettings) -> None: ... - - -class MeasurementRepository(Protocol): - def save(self, result: SweepResult, settings: AppSettings, target: SaveTarget) -> SaveArtifacts: ... - def load(self, file_path: str) -> LoadedMeasurement: ... - - -class ReferenceRepository(Protocol): - def load_reference(self, file_path: str) -> ReferenceCurve: ... +__all__ = ["MeasurementRepository", "ReferenceRepository", "SettingsRepository"] diff --git a/src/app/infrastructure/persistence/settings_defaults.py b/src/app/infrastructure/persistence/settings_defaults.py new file mode 100644 index 0000000..35e3270 --- /dev/null +++ b/src/app/infrastructure/persistence/settings_defaults.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +from app.domain.enums import ConnectionMode, CorrectionMode, CouplingMode, ImpedanceMode, MagnitudePhaseMode, TriggerMode +from app.domain.models import ( + AppSettings, + AwgSettings, + ChannelSelection, + InstrumentEndpoint, + InstrumentSetup, + OscSettings, + RunMode, + SweepSpec, +) +from app.shared.mapping import Mapping + + +class DefaultSettingsFactory: + def create(self) -> AppSettings: + return AppSettings( + schema_version=1, + freq_unit=Mapping.mapping_mhz, + sweep=SweepSpec( + start_hz=1e6, + stop_hz=100e6, + step_hz=1e6, + step_count=100, + is_log=False, + ), + run_mode=RunMode( + correction_mode=CorrectionMode.NONE, + trigger_mode=TriggerMode.FREE_RUN, + auto_range=True, + auto_reset=True, + ), + setup=InstrumentSetup( + awg=InstrumentEndpoint( + model=Mapping.mapping_DSG_4102, + connect_mode=ConnectionMode.AUTO, + visa_address="", + ip_address="0.0.0.0", + ), + osc=InstrumentEndpoint( + model=Mapping.mapping_MDO_34, + connect_mode=ConnectionMode.AUTO, + visa_address="", + ip_address="0.0.0.0", + ), + channels=ChannelSelection(awg_ch=1, osc_test_ch=1, osc_ref_ch=2, osc_trig_ch=2), + awg_settings=AwgSettings(amplitude_vpp=1.0, impedance=ImpedanceMode.R50), + osc_settings=OscSettings( + full_scale_v=1.0, + offset_v=0.0, + points=10_000, + impedance=ImpedanceMode.R50, + coupling=CouplingMode.DC, + ), + ), + magnitude_phase_mode=MagnitudePhaseMode.MAG, + auto_save_data=True, + ) diff --git a/src/app/infrastructure/persistence/settings_repo_json.py b/src/app/infrastructure/persistence/settings_repo_json.py index 8772c66..6c2e879 100644 --- a/src/app/infrastructure/persistence/settings_repo_json.py +++ b/src/app/infrastructure/persistence/settings_repo_json.py @@ -1,187 +1,43 @@ from __future__ import annotations import json -from dataclasses import asdict from pathlib import Path -from mapping import Mapping - -from app.domain.enums import ( - ConnectionMode, - CorrectionMode, - CouplingMode, - ImpedanceMode, - MagnitudePhaseMode, - TriggerMode, -) -from app.domain.models import ( - AppSettings, - AwgSettings, - ChannelSelection, - InstrumentEndpoint, - InstrumentSetup, - OscSettings, - RunMode, - SweepSpec, -) +from app.domain.models import AppSettings from app.domain.validators import validate_settings +from app.infrastructure.persistence.file_store import FileStore +from app.infrastructure.persistence.settings_defaults import DefaultSettingsFactory +from app.infrastructure.persistence.settings_serializer import SettingsSerializer +from app.runtime.paths import AppPaths class JsonSettingsRepository: - def __init__(self, config_path: Path | None = None) -> None: - root = Path(__file__).resolve().parents[4] - default_path = root / "__config__" / "settings.json" - self._path = config_path or default_path + def __init__( + self, + config_path: Path | None = None, + *, + defaults_factory: DefaultSettingsFactory | None = None, + serializer: SettingsSerializer | None = None, + file_store: FileStore | None = None, + ) -> None: + default_path = AppPaths.default().settings_path + + self._serializer = serializer or SettingsSerializer() + self._defaults_factory = defaults_factory or DefaultSettingsFactory() + self._file_store = file_store or FileStore(config_path or default_path) def load(self) -> AppSettings: - if not self._path.exists(): - settings = self._default_settings() + if not self._file_store.exists(): + settings = self._defaults_factory.create() self.save(settings) return settings - payload = json.loads(self._path.read_text(encoding="utf-8")) - settings = self._from_dict(payload) + payload = json.loads(self._file_store.read_text(encoding="utf-8")) + settings = self._serializer.from_payload(payload) validate_settings(settings) return settings def save(self, settings: AppSettings) -> None: validate_settings(settings) - self._path.parent.mkdir(parents=True, exist_ok=True) - data = self._to_dict(settings) - self._path.write_text(json.dumps(data, indent=2, ensure_ascii=True), encoding="utf-8") - - def _default_settings(self) -> AppSettings: - return AppSettings( - schema_version=1, - freq_unit=Mapping.mapping_mhz, - sweep=SweepSpec( - start_hz=1e6, - stop_hz=100e6, - step_hz=1e6, - step_count=100, - is_log=False, - ), - run_mode=RunMode( - correction_mode=CorrectionMode.NONE, - trigger_mode=TriggerMode.FREE_RUN, - auto_range=True, - auto_reset=True, - ), - setup=InstrumentSetup( - awg=InstrumentEndpoint( - model=Mapping.mapping_DSG_4102, - connect_mode=ConnectionMode.AUTO, - visa_address="", - ip_address="0.0.0.0", - ), - osc=InstrumentEndpoint( - model=Mapping.mapping_MDO_34, - connect_mode=ConnectionMode.AUTO, - visa_address="", - ip_address="0.0.0.0", - ), - channels=ChannelSelection(awg_ch=1, osc_test_ch=1, osc_ref_ch=2, osc_trig_ch=2), - awg_settings=AwgSettings(amplitude_vpp=1.0, impedance=ImpedanceMode.R50), - osc_settings=OscSettings( - full_scale_v=1.0, - offset_v=0.0, - points=10_000, - impedance=ImpedanceMode.R50, - coupling=CouplingMode.DC, - ), - ), - magnitude_phase_mode=MagnitudePhaseMode.MAG, - auto_save_data=True, - ) - - def _to_dict(self, settings: AppSettings) -> dict[str, object]: - payload = asdict(settings) - payload["run_mode"]["correction_mode"] = settings.run_mode.correction_mode.value - payload["run_mode"]["trigger_mode"] = settings.run_mode.trigger_mode.value - payload["setup"]["awg"]["connect_mode"] = settings.setup.awg.connect_mode.value - payload["setup"]["osc"]["connect_mode"] = settings.setup.osc.connect_mode.value - payload["setup"]["awg_settings"]["impedance"] = settings.setup.awg_settings.impedance.value - payload["setup"]["osc_settings"]["impedance"] = settings.setup.osc_settings.impedance.value - payload["setup"]["osc_settings"]["coupling"] = settings.setup.osc_settings.coupling.value - payload["magnitude_phase_mode"] = settings.magnitude_phase_mode.value - return payload - - def _from_dict(self, payload: dict[str, object]) -> AppSettings: - sweep_payload = payload.get("sweep", {}) - run_payload = payload.get("run_mode", {}) - setup_payload = payload.get("setup", {}) - - awg_payload = setup_payload.get("awg", {}) if isinstance(setup_payload, dict) else {} - osc_payload = setup_payload.get("osc", {}) if isinstance(setup_payload, dict) else {} - channels_payload = setup_payload.get("channels", {}) if isinstance(setup_payload, dict) else {} - awg_settings_payload = setup_payload.get("awg_settings", {}) if isinstance(setup_payload, dict) else {} - osc_settings_payload = setup_payload.get("osc_settings", {}) if isinstance(setup_payload, dict) else {} - - return AppSettings( - schema_version=int(payload.get("schema_version", 1)), - freq_unit=str(payload.get("freq_unit", Mapping.mapping_mhz)), - sweep=SweepSpec( - start_hz=float(sweep_payload.get("start_hz", 1e6)), - stop_hz=float(sweep_payload.get("stop_hz", 100e6)), - step_hz=( - None - if sweep_payload.get("step_hz") is None - else float(sweep_payload.get("step_hz", 1e6)) - ), - step_count=( - None - if sweep_payload.get("step_count") is None - else int(sweep_payload.get("step_count", 100)) - ), - is_log=bool(sweep_payload.get("is_log", False)), - ), - run_mode=RunMode( - correction_mode=CorrectionMode(str(run_payload.get("correction_mode", CorrectionMode.NONE.value))), - trigger_mode=TriggerMode(str(run_payload.get("trigger_mode", TriggerMode.FREE_RUN.value))), - auto_range=bool(run_payload.get("auto_range", True)), - auto_reset=bool(run_payload.get("auto_reset", True)), - ), - setup=InstrumentSetup( - awg=InstrumentEndpoint( - model=str(awg_payload.get("model", Mapping.mapping_DSG_4102)), - connect_mode=ConnectionMode(str(awg_payload.get("connect_mode", ConnectionMode.AUTO.value))), - visa_address=str(awg_payload.get("visa_address", "")), - ip_address=str(awg_payload.get("ip_address", "0.0.0.0")), - ), - osc=InstrumentEndpoint( - model=str(osc_payload.get("model", Mapping.mapping_MDO_34)), - connect_mode=ConnectionMode(str(osc_payload.get("connect_mode", ConnectionMode.AUTO.value))), - visa_address=str(osc_payload.get("visa_address", "")), - ip_address=str(osc_payload.get("ip_address", "0.0.0.0")), - ), - channels=ChannelSelection( - awg_ch=int(channels_payload.get("awg_ch", 1)), - osc_test_ch=int(channels_payload.get("osc_test_ch", 1)), - osc_ref_ch=( - None - if channels_payload.get("osc_ref_ch") is None - else int(channels_payload.get("osc_ref_ch")) - ), - osc_trig_ch=( - None - if channels_payload.get("osc_trig_ch") is None - else int(channels_payload.get("osc_trig_ch")) - ), - ), - awg_settings=AwgSettings( - amplitude_vpp=float(awg_settings_payload.get("amplitude_vpp", 1.0)), - impedance=ImpedanceMode(str(awg_settings_payload.get("impedance", ImpedanceMode.R50.value))), - ), - osc_settings=OscSettings( - full_scale_v=float(osc_settings_payload.get("full_scale_v", 1.0)), - offset_v=float(osc_settings_payload.get("offset_v", 0.0)), - points=int(osc_settings_payload.get("points", 10_000)), - impedance=ImpedanceMode(str(osc_settings_payload.get("impedance", ImpedanceMode.R50.value))), - coupling=CouplingMode(str(osc_settings_payload.get("coupling", CouplingMode.DC.value))), - ), - ), - magnitude_phase_mode=MagnitudePhaseMode( - str(payload.get("magnitude_phase_mode", MagnitudePhaseMode.MAG.value)) - ), - auto_save_data=bool(payload.get("auto_save_data", True)), - ) + payload = self._serializer.to_payload(settings) + self._file_store.write_text(json.dumps(payload, indent=2, ensure_ascii=True), encoding="utf-8") diff --git a/src/app/infrastructure/persistence/settings_serializer.py b/src/app/infrastructure/persistence/settings_serializer.py new file mode 100644 index 0000000..925c80b --- /dev/null +++ b/src/app/infrastructure/persistence/settings_serializer.py @@ -0,0 +1,94 @@ +from __future__ import annotations + +from dataclasses import asdict + +from app.domain.enums import ConnectionMode, CorrectionMode, CouplingMode, ImpedanceMode, MagnitudePhaseMode, TriggerMode +from app.domain.models import ( + AppSettings, + AwgSettings, + ChannelSelection, + InstrumentEndpoint, + InstrumentSetup, + OscSettings, + RunMode, + SweepSpec, +) +from app.shared.mapping import Mapping + + +class SettingsSerializer: + def to_payload(self, settings: AppSettings) -> dict[str, object]: + payload = asdict(settings) + payload["run_mode"]["correction_mode"] = settings.run_mode.correction_mode.value + payload["run_mode"]["trigger_mode"] = settings.run_mode.trigger_mode.value + payload["setup"]["awg"]["connect_mode"] = settings.setup.awg.connect_mode.value + payload["setup"]["osc"]["connect_mode"] = settings.setup.osc.connect_mode.value + payload["setup"]["awg_settings"]["impedance"] = settings.setup.awg_settings.impedance.value + payload["setup"]["osc_settings"]["impedance"] = settings.setup.osc_settings.impedance.value + payload["setup"]["osc_settings"]["coupling"] = settings.setup.osc_settings.coupling.value + payload["magnitude_phase_mode"] = settings.magnitude_phase_mode.value + return payload + + def from_payload(self, payload: dict[str, object]) -> AppSettings: + sweep_payload = payload.get("sweep", {}) + run_payload = payload.get("run_mode", {}) + setup_payload = payload.get("setup", {}) + + awg_payload = setup_payload.get("awg", {}) if isinstance(setup_payload, dict) else {} + osc_payload = setup_payload.get("osc", {}) if isinstance(setup_payload, dict) else {} + channels_payload = setup_payload.get("channels", {}) if isinstance(setup_payload, dict) else {} + awg_settings_payload = setup_payload.get("awg_settings", {}) if isinstance(setup_payload, dict) else {} + osc_settings_payload = setup_payload.get("osc_settings", {}) if isinstance(setup_payload, dict) else {} + + return AppSettings( + schema_version=int(payload.get("schema_version", 1)), + freq_unit=str(payload.get("freq_unit", Mapping.mapping_mhz)), + sweep=SweepSpec( + start_hz=float(sweep_payload.get("start_hz", 1e6)), + stop_hz=float(sweep_payload.get("stop_hz", 100e6)), + step_hz=None if sweep_payload.get("step_hz") is None else float(sweep_payload.get("step_hz", 1e6)), + step_count=None + if sweep_payload.get("step_count") is None + else int(sweep_payload.get("step_count", 100)), + is_log=bool(sweep_payload.get("is_log", False)), + ), + run_mode=RunMode( + correction_mode=CorrectionMode(str(run_payload.get("correction_mode", CorrectionMode.NONE.value))), + trigger_mode=TriggerMode(str(run_payload.get("trigger_mode", TriggerMode.FREE_RUN.value))), + auto_range=bool(run_payload.get("auto_range", True)), + auto_reset=bool(run_payload.get("auto_reset", True)), + ), + setup=InstrumentSetup( + awg=InstrumentEndpoint( + model=str(awg_payload.get("model", Mapping.mapping_DSG_4102)), + connect_mode=ConnectionMode(str(awg_payload.get("connect_mode", ConnectionMode.AUTO.value))), + visa_address=str(awg_payload.get("visa_address", "")), + ip_address=str(awg_payload.get("ip_address", "0.0.0.0")), + ), + osc=InstrumentEndpoint( + model=str(osc_payload.get("model", Mapping.mapping_MDO_34)), + connect_mode=ConnectionMode(str(osc_payload.get("connect_mode", ConnectionMode.AUTO.value))), + visa_address=str(osc_payload.get("visa_address", "")), + ip_address=str(osc_payload.get("ip_address", "0.0.0.0")), + ), + channels=ChannelSelection( + awg_ch=int(channels_payload.get("awg_ch", 1)), + osc_test_ch=int(channels_payload.get("osc_test_ch", 1)), + osc_ref_ch=None if channels_payload.get("osc_ref_ch") is None else int(channels_payload.get("osc_ref_ch")), + osc_trig_ch=None if channels_payload.get("osc_trig_ch") is None else int(channels_payload.get("osc_trig_ch")), + ), + awg_settings=AwgSettings( + amplitude_vpp=float(awg_settings_payload.get("amplitude_vpp", 1.0)), + impedance=ImpedanceMode(str(awg_settings_payload.get("impedance", ImpedanceMode.R50.value))), + ), + osc_settings=OscSettings( + full_scale_v=float(osc_settings_payload.get("full_scale_v", 1.0)), + offset_v=float(osc_settings_payload.get("offset_v", 0.0)), + points=int(osc_settings_payload.get("points", 10_000)), + impedance=ImpedanceMode(str(osc_settings_payload.get("impedance", ImpedanceMode.R50.value))), + coupling=CouplingMode(str(osc_settings_payload.get("coupling", CouplingMode.DC.value))), + ), + ), + magnitude_phase_mode=MagnitudePhaseMode(str(payload.get("magnitude_phase_mode", MagnitudePhaseMode.MAG.value))), + auto_save_data=bool(payload.get("auto_save_data", True)), + ) diff --git a/src/app/presentation/tk/app_window.py b/src/app/presentation/tk/app_window.py index 56bcc40..3a46c11 100644 --- a/src/app/presentation/tk/app_window.py +++ b/src/app/presentation/tk/app_window.py @@ -1,10 +1,8 @@ from __future__ import annotations import tkinter as tk -from tkinter import ttk - -from mapping import Mapping +from app.presentation.tk.control_panel import ControlPanel from app.presentation.tk.plot_widget import PlotWidget from app.presentation.tk.view_model import ViewModel @@ -28,7 +26,9 @@ def __init__(self, vm: ViewModel | None = None) -> None: right = tk.Frame(container) right.pack(side=tk.LEFT, fill=tk.BOTH, expand=True, padx=8, pady=8) - self._build_controls(left) + self.control_panel = ControlPanel(left, self.vm) + self.control_panel.pack(fill=tk.Y) + self._alias_control_widgets() self.plot_widget = PlotWidget(right) self.plot_widget.frame.pack(fill=tk.BOTH, expand=True) @@ -49,196 +49,6 @@ def __init__(self, vm: ViewModel | None = None) -> None: self.osc_light = self.canvas_osc.create_oval(2, 2, 12, 12, fill="red") tk.Label(status_bar, text="OSC").pack(side=tk.RIGHT) - def _build_controls(self, parent: tk.Misc) -> None: - row = 0 - - def add_label(text: str, r: int, c: int = 0) -> None: - tk.Label(parent, text=text).grid(row=r, column=c, sticky="w", padx=3, pady=2) - - add_label("AWG model", row) - ttk.Combobox(parent, textvariable=self.vm.awg_model, values=Mapping.values_awg, width=12).grid( - row=row, column=1, sticky="ew" - ) - row += 1 - - add_label("OSC model", row) - ttk.Combobox(parent, textvariable=self.vm.osc_model, values=Mapping.values_osc, width=12).grid( - row=row, column=1, sticky="ew" - ) - row += 1 - - add_label("AWG conn", row) - ttk.Combobox(parent, textvariable=self.vm.awg_connect_mode, values=["auto", "lan"], width=12).grid( - row=row, column=1, sticky="ew" - ) - row += 1 - - add_label("AWG VISA", row) - tk.Entry(parent, textvariable=self.vm.awg_visa, width=28).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("AWG IP", row) - tk.Entry(parent, textvariable=self.vm.awg_ip, width=28).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC conn", row) - ttk.Combobox(parent, textvariable=self.vm.osc_connect_mode, values=["auto", "lan"], width=12).grid( - row=row, column=1, sticky="ew" - ) - row += 1 - - add_label("OSC VISA", row) - tk.Entry(parent, textvariable=self.vm.osc_visa, width=28).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC IP", row) - tk.Entry(parent, textvariable=self.vm.osc_ip, width=28).grid(row=row, column=1, sticky="ew") - row += 1 - - ttk.Separator(parent, orient=tk.HORIZONTAL).grid(row=row, column=0, columnspan=2, sticky="ew", pady=8) - row += 1 - - add_label("Freq unit", row) - ttk.Combobox(parent, textvariable=self.vm.freq_unit, values=Mapping.values_freq_unit, width=8).grid( - row=row, column=1, sticky="ew" - ) - row += 1 - - add_label("Start", row) - tk.Entry(parent, textvariable=self.vm.start_freq).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("Stop", row) - tk.Entry(parent, textvariable=self.vm.stop_freq).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("Step", row) - tk.Entry(parent, textvariable=self.vm.step_freq).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("Step count", row) - tk.Entry(parent, textvariable=self.vm.step_count).grid(row=row, column=1, sticky="ew") - row += 1 - - tk.Checkbutton(parent, text="Log sweep", variable=self.vm.is_log).grid(row=row, column=0, columnspan=2, sticky="w") - row += 1 - - ttk.Separator(parent, orient=tk.HORIZONTAL).grid(row=row, column=0, columnspan=2, sticky="ew", pady=8) - row += 1 - - add_label("AWG amp (Vpp)", row) - tk.Entry(parent, textvariable=self.vm.awg_amp).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("AWG imp", row) - ttk.Combobox(parent, textvariable=self.vm.awg_imp, values=["50", "INF"], width=8).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC range", row) - tk.Entry(parent, textvariable=self.vm.osc_range).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC offset", row) - tk.Entry(parent, textvariable=self.vm.osc_offset).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC points", row) - tk.Entry(parent, textvariable=self.vm.osc_points).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC imp", row) - ttk.Combobox(parent, textvariable=self.vm.osc_imp, values=["50", "INF"], width=8).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC coup", row) - ttk.Combobox(parent, textvariable=self.vm.osc_coupling, values=["DC", "AC"], width=8).grid(row=row, column=1, sticky="ew") - row += 1 - - ttk.Separator(parent, orient=tk.HORIZONTAL).grid(row=row, column=0, columnspan=2, sticky="ew", pady=8) - row += 1 - - add_label("AWG ch", row) - tk.Entry(parent, textvariable=self.vm.awg_ch).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC test ch", row) - tk.Entry(parent, textvariable=self.vm.osc_test_ch).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC ref ch", row) - tk.Entry(parent, textvariable=self.vm.osc_ref_ch).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("OSC trig ch", row) - tk.Entry(parent, textvariable=self.vm.osc_trig_ch).grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("Correction", row) - ttk.Combobox(parent, textvariable=self.vm.correction_mode, values=["none", "single", "dual"], width=10).grid( - row=row, column=1, sticky="ew" - ) - row += 1 - - add_label("Trigger", row) - ttk.Combobox(parent, textvariable=self.vm.trigger_mode, values=["free_run", "triggered"], width=10).grid( - row=row, column=1, sticky="ew" - ) - row += 1 - - tk.Checkbutton(parent, text="Auto range", variable=self.vm.auto_range).grid(row=row, column=0, columnspan=2, sticky="w") - row += 1 - tk.Checkbutton(parent, text="Auto reset", variable=self.vm.auto_reset).grid(row=row, column=0, columnspan=2, sticky="w") - row += 1 - tk.Checkbutton(parent, text="Enable calibration", variable=self.vm.calibration_enabled).grid( - row=row, column=0, columnspan=2, sticky="w" - ) - row += 1 - tk.Checkbutton(parent, text="Auto save data", variable=self.vm.auto_save_data).grid( - row=row, column=0, columnspan=2, sticky="w" - ) - row += 1 - - add_label("Figure", row) - self.cmb_figure = ttk.Combobox(parent, textvariable=self.vm.figure_mode, values=["gain", "gain_db"], width=10) - self.cmb_figure.grid(row=row, column=1, sticky="ew") - row += 1 - - add_label("Display", row) - self.cmb_mag_phase = ttk.Combobox( - parent, - textvariable=self.vm.magnitude_phase_mode, - values=["magnitude", "phase", "magnitude_phase"], - width=14, - ) - self.cmb_mag_phase.grid(row=row, column=1, sticky="ew") - row += 1 - - buttons = tk.Frame(parent) - buttons.grid(row=row, column=0, columnspan=2, sticky="ew", pady=10) - - self.btn_start = tk.Button(buttons, text="Start", width=9) - self.btn_start.pack(side=tk.LEFT, padx=2) - self.btn_stop = tk.Button(buttons, text="Stop", width=9, state="disabled") - self.btn_stop.pack(side=tk.LEFT, padx=2) - - self.btn_save_data = tk.Button(buttons, text="Save Data", width=9) - self.btn_save_data.pack(side=tk.LEFT, padx=2) - self.btn_load_data = tk.Button(buttons, text="Load Data", width=9) - self.btn_load_data.pack(side=tk.LEFT, padx=2) - - row += 1 - buttons2 = tk.Frame(parent) - buttons2.grid(row=row, column=0, columnspan=2, sticky="ew", pady=5) - - self.btn_load_ref = tk.Button(buttons2, text="Load Ref", width=9) - self.btn_load_ref.pack(side=tk.LEFT, padx=2) - self.btn_save_settings = tk.Button(buttons2, text="Save Settings", width=11) - self.btn_save_settings.pack(side=tk.LEFT, padx=2) - self.btn_load_settings = tk.Button(buttons2, text="Load Settings", width=11) - self.btn_load_settings.pack(side=tk.LEFT, padx=2) - - parent.grid_columnconfigure(1, weight=1) - def bind_actions( self, *, @@ -253,16 +63,17 @@ def bind_actions( on_figure_change, on_mag_phase_change, ) -> None: - self.btn_start.configure(command=on_start) - self.btn_stop.configure(command=on_stop) - self.btn_save_data.configure(command=on_save_data) - self.btn_load_data.configure(command=on_load_data) - self.btn_load_ref.configure(command=on_load_ref) - self.btn_save_settings.configure(command=on_save_settings) - self.btn_load_settings.configure(command=on_load_settings) - - self.cmb_figure.bind("<>", lambda _e: on_figure_change()) - self.cmb_mag_phase.bind("<>", lambda _e: on_mag_phase_change()) + self.control_panel.bind_actions( + on_start=on_start, + on_stop=on_stop, + on_save_data=on_save_data, + on_load_data=on_load_data, + on_load_ref=on_load_ref, + on_save_settings=on_save_settings, + on_load_settings=on_load_settings, + on_figure_change=on_figure_change, + on_mag_phase_change=on_mag_phase_change, + ) self._on_close = on_close def set_connection_status(self, awg_connected: bool, osc_connected: bool) -> None: @@ -274,3 +85,14 @@ def on_close(self) -> None: self._on_close() else: self.destroy() + + def _alias_control_widgets(self) -> None: + self.btn_start = self.control_panel.btn_start + self.btn_stop = self.control_panel.btn_stop + self.btn_save_data = self.control_panel.btn_save_data + self.btn_load_data = self.control_panel.btn_load_data + self.btn_load_ref = self.control_panel.btn_load_ref + self.btn_save_settings = self.control_panel.btn_save_settings + self.btn_load_settings = self.control_panel.btn_load_settings + self.cmb_figure = self.control_panel.cmb_figure + self.cmb_mag_phase = self.control_panel.cmb_mag_phase diff --git a/src/app/presentation/tk/control_panel.py b/src/app/presentation/tk/control_panel.py new file mode 100644 index 0000000..510ed55 --- /dev/null +++ b/src/app/presentation/tk/control_panel.py @@ -0,0 +1,222 @@ +from __future__ import annotations + +import tkinter as tk +from tkinter import ttk + +from app.presentation.tk.view_model import ViewModel +from app.shared.mapping import Mapping + + +class ControlPanel(tk.Frame): + def __init__(self, parent: tk.Misc, vm: ViewModel) -> None: + super().__init__(parent) + self._vm = vm + self._build() + + def bind_actions( + self, + *, + on_start, + on_stop, + on_save_data, + on_load_data, + on_load_ref, + on_save_settings, + on_load_settings, + on_figure_change, + on_mag_phase_change, + ) -> None: + self.btn_start.configure(command=on_start) + self.btn_stop.configure(command=on_stop) + self.btn_save_data.configure(command=on_save_data) + self.btn_load_data.configure(command=on_load_data) + self.btn_load_ref.configure(command=on_load_ref) + self.btn_save_settings.configure(command=on_save_settings) + self.btn_load_settings.configure(command=on_load_settings) + + self.cmb_figure.bind("<>", lambda _e: on_figure_change()) + self.cmb_mag_phase.bind("<>", lambda _e: on_mag_phase_change()) + + def _build(self) -> None: + row = 0 + + def add_label(text: str, r: int, c: int = 0) -> None: + tk.Label(self, text=text).grid(row=r, column=c, sticky="w", padx=3, pady=2) + + add_label("AWG model", row) + ttk.Combobox(self, textvariable=self._vm.awg_model, values=Mapping.values_awg, width=12).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + add_label("OSC model", row) + ttk.Combobox(self, textvariable=self._vm.osc_model, values=Mapping.values_osc, width=12).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + add_label("AWG conn", row) + ttk.Combobox(self, textvariable=self._vm.awg_connect_mode, values=["auto", "lan"], width=12).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + add_label("AWG VISA", row) + tk.Entry(self, textvariable=self._vm.awg_visa, width=28).grid(row=row, column=1, sticky="ew") + row += 1 + + add_label("AWG IP", row) + tk.Entry(self, textvariable=self._vm.awg_ip, width=28).grid(row=row, column=1, sticky="ew") + row += 1 + + add_label("OSC conn", row) + ttk.Combobox(self, textvariable=self._vm.osc_connect_mode, values=["auto", "lan"], width=12).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + add_label("OSC VISA", row) + tk.Entry(self, textvariable=self._vm.osc_visa, width=28).grid(row=row, column=1, sticky="ew") + row += 1 + + add_label("OSC IP", row) + tk.Entry(self, textvariable=self._vm.osc_ip, width=28).grid(row=row, column=1, sticky="ew") + row += 1 + + self._separator(row) + row += 1 + + add_label("Freq unit", row) + ttk.Combobox(self, textvariable=self._vm.freq_unit, values=Mapping.values_freq_unit, width=8).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + for label, variable in ( + ("Start", self._vm.start_freq), + ("Stop", self._vm.stop_freq), + ("Step", self._vm.step_freq), + ("Step count", self._vm.step_count), + ): + add_label(label, row) + tk.Entry(self, textvariable=variable).grid(row=row, column=1, sticky="ew") + row += 1 + + tk.Checkbutton(self, text="Log sweep", variable=self._vm.is_log).grid( + row=row, column=0, columnspan=2, sticky="w" + ) + row += 1 + + self._separator(row) + row += 1 + + for label, variable in ( + ("AWG amp (Vpp)", self._vm.awg_amp), + ("OSC range", self._vm.osc_range), + ("OSC offset", self._vm.osc_offset), + ("OSC points", self._vm.osc_points), + ): + if label == "OSC range": + add_label("AWG imp", row) + ttk.Combobox(self, textvariable=self._vm.awg_imp, values=["50", "INF"], width=8).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + add_label(label, row) + tk.Entry(self, textvariable=variable).grid(row=row, column=1, sticky="ew") + row += 1 + + add_label("OSC imp", row) + ttk.Combobox(self, textvariable=self._vm.osc_imp, values=["50", "INF"], width=8).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + add_label("OSC coup", row) + ttk.Combobox(self, textvariable=self._vm.osc_coupling, values=["DC", "AC"], width=8).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + self._separator(row) + row += 1 + + for label, variable in ( + ("AWG ch", self._vm.awg_ch), + ("OSC test ch", self._vm.osc_test_ch), + ("OSC ref ch", self._vm.osc_ref_ch), + ("OSC trig ch", self._vm.osc_trig_ch), + ): + add_label(label, row) + tk.Entry(self, textvariable=variable).grid(row=row, column=1, sticky="ew") + row += 1 + + add_label("Correction", row) + ttk.Combobox(self, textvariable=self._vm.correction_mode, values=["none", "single", "dual"], width=10).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + add_label("Trigger", row) + ttk.Combobox(self, textvariable=self._vm.trigger_mode, values=["free_run", "triggered"], width=10).grid( + row=row, column=1, sticky="ew" + ) + row += 1 + + for label, variable in ( + ("Auto range", self._vm.auto_range), + ("Auto reset", self._vm.auto_reset), + ("Enable calibration", self._vm.calibration_enabled), + ("Auto save data", self._vm.auto_save_data), + ): + tk.Checkbutton(self, text=label, variable=variable).grid(row=row, column=0, columnspan=2, sticky="w") + row += 1 + + add_label("Figure", row) + self.cmb_figure = ttk.Combobox(self, textvariable=self._vm.figure_mode, values=["gain", "gain_db"], width=10) + self.cmb_figure.grid(row=row, column=1, sticky="ew") + row += 1 + + add_label("Display", row) + self.cmb_mag_phase = ttk.Combobox( + self, + textvariable=self._vm.magnitude_phase_mode, + values=["magnitude", "phase", "magnitude_phase"], + width=14, + ) + self.cmb_mag_phase.grid(row=row, column=1, sticky="ew") + row += 1 + + row = self._build_primary_buttons(row) + self._build_secondary_buttons(row) + self.grid_columnconfigure(1, weight=1) + + def _separator(self, row: int) -> None: + ttk.Separator(self, orient=tk.HORIZONTAL).grid(row=row, column=0, columnspan=2, sticky="ew", pady=8) + + def _build_primary_buttons(self, row: int) -> int: + buttons = tk.Frame(self) + buttons.grid(row=row, column=0, columnspan=2, sticky="ew", pady=10) + + self.btn_start = tk.Button(buttons, text="Start", width=9) + self.btn_start.pack(side=tk.LEFT, padx=2) + self.btn_stop = tk.Button(buttons, text="Stop", width=9, state="disabled") + self.btn_stop.pack(side=tk.LEFT, padx=2) + + self.btn_save_data = tk.Button(buttons, text="Save Data", width=9) + self.btn_save_data.pack(side=tk.LEFT, padx=2) + self.btn_load_data = tk.Button(buttons, text="Load Data", width=9) + self.btn_load_data.pack(side=tk.LEFT, padx=2) + return row + 1 + + def _build_secondary_buttons(self, row: int) -> None: + buttons = tk.Frame(self) + buttons.grid(row=row, column=0, columnspan=2, sticky="ew", pady=5) + + self.btn_load_ref = tk.Button(buttons, text="Load Ref", width=9) + self.btn_load_ref.pack(side=tk.LEFT, padx=2) + self.btn_save_settings = tk.Button(buttons, text="Save Settings", width=11) + self.btn_save_settings.pack(side=tk.LEFT, padx=2) + self.btn_load_settings = tk.Button(buttons, text="Load Settings", width=11) + self.btn_load_settings.pack(side=tk.LEFT, padx=2) + diff --git a/src/app/presentation/tk/controller.py b/src/app/presentation/tk/controller.py index 536734c..6684e52 100644 --- a/src/app/presentation/tk/controller.py +++ b/src/app/presentation/tk/controller.py @@ -2,36 +2,26 @@ import queue import threading -from pathlib import Path - -from app.application.dto import SaveTarget, StartSweepCommand -from app.application.events import ( - ConnectionStatusUpdated, - SweepCompleted, - SweepDataUpdated, - SweepFailed, - SweepProgress, - SweepStarted, - SweepStopped, - SweepWarning, -) +from collections.abc import Callable + +from app.application.events import EventEmitter +from app.application.services.sweep_task_runner import SweepTaskRunner from app.application.services.connection_monitor import ConnectionMonitor from app.application.use_cases.load_measurement import LoadMeasurementUseCase from app.application.use_cases.load_reference import LoadReferenceUseCase from app.application.use_cases.save_measurement import SaveMeasurementUseCase from app.application.use_cases.settings_use_case import SettingsUseCase -from app.application.use_cases.start_sweep import StartSweepUseCase -from app.application.use_cases.stop_sweep import StopSweepUseCase -from app.domain.models import AppSettings, SweepResult -from app.infrastructure.instruments.equips_factory import create_instrument_ports, resolve_visa_address -from app.infrastructure.instruments.ports import ResourceScannerPort +from app.application.ports.instruments import InstrumentPortsFactory, ResourceScannerPort +from app.domain.models import InstrumentEndpoint from app.presentation.tk import dialogs from app.presentation.tk.app_window import AppWindow from app.presentation.tk.mapper import settings_to_vm, vm_to_settings +from app.presentation.tk.ui_event_handler import UiEventHandler from app.presentation.tk.view_model import ViewModel +from app.runtime.paths import AppPaths -class TkController: +class TkController(EventEmitter): def __init__( self, *, @@ -42,28 +32,37 @@ def __init__( load_measurement_use_case: LoadMeasurementUseCase, load_reference_use_case: LoadReferenceUseCase, scanner: ResourceScannerPort, + ports_factory: InstrumentPortsFactory, + resolve_address: Callable[[InstrumentEndpoint], str], + paths: AppPaths | None = None, ) -> None: self.window = window self.vm = vm - self.settings_use_case = settings_use_case self.save_measurement_use_case = save_measurement_use_case self.load_measurement_use_case = load_measurement_use_case self.load_reference_use_case = load_reference_use_case self._event_queue: queue.Queue[object] = queue.Queue() - self._latest_result = SweepResult() self._reference_interpolator = None - - self._ports = None - self._sweep_thread: threading.Thread | None = None - self._stop_use_case: StopSweepUseCase | None = None - - self._root_dir = Path(__file__).resolve().parents[4] + self._paths = paths or AppPaths.default() + self._resolve_address = resolve_address + self._connection_target_lock = threading.Lock() + self._awg_target_address = "" + self._osc_target_address = "" + self._closing = False + + self._ui_handler = UiEventHandler(window=window, vm=vm) + self._task_runner = SweepTaskRunner( + emitter=self, + save_measurement_use_case=save_measurement_use_case, + auto_save_dir=self._paths.measurement_dir, + ports_factory=ports_factory, + ) self._monitor = ConnectionMonitor( scanner=scanner, - get_awg_address=self._get_awg_target_address, - get_osc_address=self._get_osc_target_address, + get_awg_address=self._get_cached_awg_target_address, + get_osc_address=self._get_cached_osc_target_address, emitter=self, ) @@ -87,6 +86,7 @@ def initialize(self) -> None: except Exception as exc: # noqa: BLE001 dialogs.show_warning(self.window, f"Failed to load settings: {exc}") + self._refresh_connection_targets() self._monitor.start() self.window.after(100, self._process_events) self.on_figure_change() @@ -96,43 +96,29 @@ def emit(self, event: object) -> None: self._event_queue.put(event) def on_start(self) -> None: - if self._sweep_thread and self._sweep_thread.is_alive(): + if self._task_runner.is_running(): return try: settings = vm_to_settings(self.vm) - ports = create_instrument_ports(settings.setup) + self._refresh_connection_targets(settings) + self._task_runner.start( + settings=settings, + calibration_enabled=bool(self.vm.calibration_enabled.get()), + reference_interpolator=self._reference_interpolator, + ) + self._ui_handler.prepare_for_sweep_start() except Exception as exc: # noqa: BLE001 dialogs.show_warning(self.window, f"Invalid settings: {exc}") - return - - stop_event = threading.Event() - self._stop_use_case = StopSweepUseCase(stop_event=stop_event) - - cmd = StartSweepCommand( - settings=settings, - calibration_enabled=bool(self.vm.calibration_enabled.get()), - reference_interpolator=self._reference_interpolator, - ) - - self._ports = ports - start_use_case = StartSweepUseCase(awg=ports.awg, osc=ports.osc, stop_event=stop_event) - - self.window.btn_start.configure(state="disabled") - self.window.btn_stop.configure(state="normal") - self.vm.status_text.set("Sweep started") - - self._sweep_thread = threading.Thread(target=self._run_sweep, args=(start_use_case, cmd), daemon=True) - self._sweep_thread.start() def on_stop(self) -> None: - if self._stop_use_case is not None: - self._stop_use_case.stop() + self._task_runner.stop() def on_save_settings(self) -> None: try: settings = vm_to_settings(self.vm) self.settings_use_case.save(settings) + self._refresh_connection_targets(settings) dialogs.show_info(self.window, "Settings saved") except Exception as exc: # noqa: BLE001 dialogs.show_warning(self.window, f"Failed to save settings: {exc}") @@ -141,6 +127,7 @@ def on_load_settings(self) -> None: try: settings = self.settings_use_case.load() settings_to_vm(settings, self.vm) + self._refresh_connection_targets(settings) self.on_figure_change() self.on_mag_phase_change() dialogs.show_info(self.window, "Settings loaded") @@ -148,15 +135,15 @@ def on_load_settings(self) -> None: dialogs.show_warning(self.window, f"Failed to load settings: {exc}") def on_save_data(self) -> None: - if self._latest_result.is_empty: + if self._ui_handler.latest_result.is_empty: dialogs.show_warning(self.window, "No measurement data available") return fp = dialogs.ask_save_file( title="Save measurement", - initial_dir=self._root_dir / "__data__", + initial_dir=self._paths.data_dir, initial_name="measurement", - filetypes=[("All files", "*.*")], + filetypes=[("MAT files", "*.mat"), ("All files", "*.*")], ) if fp is None: return @@ -164,9 +151,9 @@ def on_save_data(self) -> None: try: settings = vm_to_settings(self.vm) artifacts = self.save_measurement_use_case.execute( - result=self._latest_result, + result=self._ui_handler.latest_result, settings=settings, - target=SaveTarget(base_path=fp, include_timestamp=False, figures=self.window.plot_widget.figures()), + target=dialogs_to_target(fp, self.window), ) dialogs.show_info(self.window, f"Saved: {artifacts.mat_path.name}") except Exception as exc: # noqa: BLE001 @@ -175,7 +162,7 @@ def on_save_data(self) -> None: def on_load_data(self) -> None: fp = dialogs.ask_open_file( title="Load measurement", - initial_dir=self._root_dir / "__data__", + initial_dir=self._paths.data_dir, filetypes=[("Measurement", "*.mat *.csv"), ("All files", "*.*")], ) if fp is None: @@ -183,12 +170,7 @@ def on_load_data(self) -> None: try: loaded = self.load_measurement_use_case.execute(str(fp)) - self._latest_result = loaded.result - self.window.plot_widget.update_result( - self._latest_result, - self.vm.freq_unit.get(), - self.vm.magnitude_phase_mode.get(), - ) + self._ui_handler.set_result(loaded.result) dialogs.show_info(self.window, "Measurement loaded") except Exception as exc: # noqa: BLE001 dialogs.show_warning(self.window, f"Failed to load data: {exc}") @@ -196,7 +178,7 @@ def on_load_data(self) -> None: def on_load_reference(self) -> None: fp = dialogs.ask_open_file( title="Load reference", - initial_dir=self._root_dir / "__data__", + initial_dir=self._paths.data_dir, filetypes=[("MAT files", "*.mat"), ("All files", "*.*")], ) if fp is None: @@ -214,135 +196,62 @@ def on_figure_change(self) -> None: self.window.plot_widget.set_mode(self.vm.figure_mode.get()) def on_mag_phase_change(self) -> None: - self.window.plot_widget.update_result( - self._latest_result, - self.vm.freq_unit.get(), - self.vm.magnitude_phase_mode.get(), - ) + self._ui_handler.refresh_plot() def on_close(self) -> None: + self._closing = True self._monitor.stop() - if self._stop_use_case is not None: - self._stop_use_case.stop() - try: settings = vm_to_settings(self.vm) self.settings_use_case.save(settings) except Exception: pass - self._close_ports() + self._task_runner.shutdown() + self._drain_event_queue() self.window.destroy() - def _run_sweep(self, start_use_case: StartSweepUseCase, cmd: StartSweepCommand) -> None: - try: - result = start_use_case.run(cmd, self) - if not result.is_empty: - self._latest_result = result - - if self.vm.auto_save_data.get(): - settings = cmd.settings - target = SaveTarget( - base_path=self._root_dir / "__data__" / "measurement", - include_timestamp=True, - figures={}, - ) - self.save_measurement_use_case.execute(result=result, settings=settings, target=target) - except Exception as exc: # noqa: BLE001 - self.emit(SweepFailed(error_code="SWEEP_THREAD", message=str(exc))) - finally: - self._close_ports() - - def _close_ports(self) -> None: - if self._ports is None: - return - try: - self._ports.awg.close() - except Exception: - pass - try: - self._ports.osc.close() - except Exception: - pass - self._ports = None - def _process_events(self) -> None: + self._drain_event_queue() + if not self._closing: + self._refresh_connection_targets() + self.window.after(100, self._process_events) + + def _drain_event_queue(self) -> None: try: while True: event = self._event_queue.get_nowait() - self._handle_event(event) + self._ui_handler.handle(event) except queue.Empty: pass - finally: - self.window.after(100, self._process_events) - - def _handle_event(self, event: object) -> None: - if isinstance(event, ConnectionStatusUpdated): - self.window.set_connection_status(event.awg_connected, event.osc_connected) - return - if isinstance(event, SweepStarted): - self.vm.status_text.set(f"Sweep started ({event.total_points} points)") - return - - if isinstance(event, SweepProgress): - self.vm.status_text.set( - f"Freq {event.freq_hz:.2f} Hz ({event.point_index}/{event.total_points})" - ) - return - - if isinstance(event, SweepDataUpdated): - self._latest_result = event.partial_result - self.window.plot_widget.update_result( - self._latest_result, - self.vm.freq_unit.get(), - self.vm.magnitude_phase_mode.get(), - ) - return + def _refresh_connection_targets(self, settings=None) -> None: + try: + settings = settings or vm_to_settings(self.vm) + awg_address = self._resolve_address(settings.setup.awg) + osc_address = self._resolve_address(settings.setup.osc) + except Exception: + awg_address = "" + osc_address = "" - if isinstance(event, SweepWarning): - if event.code in {"READY", "FREQ_MISMATCH", "AMP_MISMATCH"}: - self.vm.status_text.set(event.message) - else: - dialogs.show_warning(self.window, event.message) - return + with self._connection_target_lock: + self._awg_target_address = awg_address + self._osc_target_address = osc_address - if isinstance(event, SweepFailed): - self.vm.status_text.set(f"Sweep failed: {event.message}") - self.window.btn_start.configure(state="normal") - self.window.btn_stop.configure(state="disabled") - dialogs.show_warning(self.window, event.message) - return + def _get_cached_awg_target_address(self) -> str: + with self._connection_target_lock: + return self._awg_target_address - if isinstance(event, SweepStopped): - self._latest_result = event.result - self.vm.status_text.set("Sweep stopped") - self.window.btn_start.configure(state="normal") - self.window.btn_stop.configure(state="disabled") - return + def _get_cached_osc_target_address(self) -> str: + with self._connection_target_lock: + return self._osc_target_address - if isinstance(event, SweepCompleted): - self._latest_result = event.result - self.vm.status_text.set("Sweep completed") - self.window.plot_widget.update_result( - self._latest_result, - self.vm.freq_unit.get(), - self.vm.magnitude_phase_mode.get(), - ) - self.window.btn_start.configure(state="normal") - self.window.btn_stop.configure(state="disabled") - return - def _get_awg_target_address(self) -> str: - try: - settings = vm_to_settings(self.vm) - return resolve_visa_address(settings.setup.awg) - except Exception: - return "" +def dialogs_to_target(path, window: AppWindow): + from app.application.dto import SaveTarget - def _get_osc_target_address(self) -> str: - try: - settings = vm_to_settings(self.vm) - return resolve_visa_address(settings.setup.osc) - except Exception: - return "" + return SaveTarget( + base_path=path, + include_timestamp=False, + figures=window.plot_widget.figures(), + ) diff --git a/src/app/presentation/tk/mapper.py b/src/app/presentation/tk/mapper.py index 14c5714..c67834e 100644 --- a/src/app/presentation/tk/mapper.py +++ b/src/app/presentation/tk/mapper.py @@ -1,6 +1,6 @@ from __future__ import annotations -from cvtTools import CvtTools +from app.shared.cvt_tools import CvtTools from app.domain.enums import ( ConnectionMode, diff --git a/src/app/presentation/tk/plot_widget.py b/src/app/presentation/tk/plot_widget.py index 44c9de1..271412c 100644 --- a/src/app/presentation/tk/plot_widget.py +++ b/src/app/presentation/tk/plot_widget.py @@ -6,8 +6,8 @@ from matplotlib.backends.backend_tkagg import FigureCanvasTkAgg from matplotlib.figure import Figure -from cvtTools import CvtTools -from mapping import Mapping +from app.shared.cvt_tools import CvtTools +from app.shared.mapping import Mapping from app.domain.models import SweepResult diff --git a/src/app/presentation/tk/sweep_task_runner.py b/src/app/presentation/tk/sweep_task_runner.py new file mode 100644 index 0000000..8c65208 --- /dev/null +++ b/src/app/presentation/tk/sweep_task_runner.py @@ -0,0 +1,5 @@ +from __future__ import annotations + +from app.application.services.sweep_task_runner import SweepTaskRunner + +__all__ = ["SweepTaskRunner"] diff --git a/src/app/presentation/tk/ui_event_handler.py b/src/app/presentation/tk/ui_event_handler.py new file mode 100644 index 0000000..0330dea --- /dev/null +++ b/src/app/presentation/tk/ui_event_handler.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +from app.application.events import ( + ConnectionStatusUpdated, + SweepCompleted, + SweepDataUpdated, + SweepFailed, + SweepProgress, + SweepStarted, + SweepStopped, + SweepWarning, +) +from app.domain.models import SweepResult +from app.presentation.tk import dialogs +from app.presentation.tk.app_window import AppWindow +from app.presentation.tk.view_model import ViewModel + + +class UiEventHandler: + def __init__(self, *, window: AppWindow, vm: ViewModel) -> None: + self._window = window + self._vm = vm + self._latest_result = SweepResult() + + @property + def latest_result(self) -> SweepResult: + return self._latest_result + + def prepare_for_sweep_start(self) -> None: + self._window.btn_start.configure(state="disabled") + self._window.btn_stop.configure(state="normal") + self._vm.status_text.set("Sweep started") + + def set_result(self, result: SweepResult, *, refresh_plot: bool = True) -> None: + self._latest_result = result + if refresh_plot: + self.refresh_plot() + + def refresh_plot(self) -> None: + self._window.plot_widget.update_result( + self._latest_result, + self._vm.freq_unit.get(), + self._vm.magnitude_phase_mode.get(), + ) + + def handle(self, event: object) -> None: + if isinstance(event, ConnectionStatusUpdated): + self._window.set_connection_status(event.awg_connected, event.osc_connected) + return + + if isinstance(event, SweepStarted): + self._vm.status_text.set(f"Sweep started ({event.total_points} points)") + return + + if isinstance(event, SweepProgress): + self._vm.status_text.set(f"Freq {event.freq_hz:.2f} Hz ({event.point_index}/{event.total_points})") + return + + if isinstance(event, SweepDataUpdated): + self.set_result(event.partial_result) + return + + if isinstance(event, SweepWarning): + if event.code in {"READY", "FREQ_MISMATCH", "AMP_MISMATCH"}: + self._vm.status_text.set(event.message) + else: + dialogs.show_warning(self._window, event.message) + return + + if isinstance(event, SweepFailed): + self._vm.status_text.set(f"Sweep failed: {event.message}") + self._window.btn_start.configure(state="normal") + self._window.btn_stop.configure(state="disabled") + dialogs.show_warning(self._window, event.message) + return + + if isinstance(event, SweepStopped): + self.set_result(event.result, refresh_plot=False) + self._vm.status_text.set("Sweep stopped") + self._window.btn_start.configure(state="normal") + self._window.btn_stop.configure(state="disabled") + return + + if isinstance(event, SweepCompleted): + self.set_result(event.result) + self._vm.status_text.set("Sweep completed") + self._window.btn_start.configure(state="normal") + self._window.btn_stop.configure(state="disabled") diff --git a/src/app/presentation/tk/view_model.py b/src/app/presentation/tk/view_model.py index 5710992..0ee8146 100644 --- a/src/app/presentation/tk/view_model.py +++ b/src/app/presentation/tk/view_model.py @@ -2,7 +2,7 @@ import tkinter as tk -from mapping import Mapping +from app.shared.mapping import Mapping class ViewModel: diff --git a/src/app/runtime/__init__.py b/src/app/runtime/__init__.py new file mode 100644 index 0000000..fdc032d --- /dev/null +++ b/src/app/runtime/__init__.py @@ -0,0 +1,2 @@ +"""Runtime composition helpers for the desktop application.""" + diff --git a/src/app/runtime/paths.py b/src/app/runtime/paths.py new file mode 100644 index 0000000..7e107cc --- /dev/null +++ b/src/app/runtime/paths.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from dataclasses import dataclass +import os +from pathlib import Path + + +APP_ROOT_ENV = "AUTO_LOAD_OFF_TEST_ROOT" + + +@dataclass(frozen=True, slots=True) +class AppPaths: + root_dir: Path + config_dir: Path + data_dir: Path + + @classmethod + def from_root(cls, root_dir: Path) -> "AppPaths": + root = root_dir.resolve() + return cls( + root_dir=root, + config_dir=root / "__config__", + data_dir=root / "__data__", + ) + + @classmethod + def default(cls) -> "AppPaths": + configured_root = os.environ.get(APP_ROOT_ENV) + if configured_root: + return cls.from_root(Path(configured_root)) + return cls.from_root(Path.cwd()) + + @property + def settings_path(self) -> Path: + return self.config_dir / "settings.json" + + @property + def measurement_dir(self) -> Path: + return self.data_dir / "measurement" diff --git a/src/app/shared/__init__.py b/src/app/shared/__init__.py new file mode 100644 index 0000000..7cb28b5 --- /dev/null +++ b/src/app/shared/__init__.py @@ -0,0 +1,4 @@ +from app.shared.cvt_tools import CvtTools +from app.shared.mapping import Mapping + +__all__ = ["CvtTools", "Mapping"] diff --git a/src/app/shared/cvt_tools.py b/src/app/shared/cvt_tools.py new file mode 100644 index 0000000..2585b17 --- /dev/null +++ b/src/app/shared/cvt_tools.py @@ -0,0 +1,126 @@ +from __future__ import annotations + +import math +import re + +import numpy as np + + +class CvtTools: + @staticmethod + def parse_general_val(input: str, default_unit: str | None = None) -> float | int: + input = input.replace(" ", "") + input_match = re.search(r"([+-]?\d*(?:\.\d+)?(?:[eE][+-]?\d+)?)([A-Za-zµ]?)", input) + if input_match is None: + return 0.0 + + input_val = input_match.group(1) + input_unit = input_match.group(2) + + try: + input_val = int(input_val) if input_val.isdigit() else float(input_val) + except Exception: + return 0.0 + + if not input_unit: + val = input_val * CvtTools.convert_general_unit(default_unit) + try: + return int(val) if float(val).is_integer() else val + except Exception: + return val + + prefix = input_unit[0] + scale = CvtTools._scale_for_prefix(prefix) + val = input_val * scale + try: + return int(val) if float(val).is_integer() else val + except Exception: + return val + + @staticmethod + def convert_general_unit(unit: str | None) -> float | int: + if not unit: + return 1 + + input_match = re.search(r"([+-]?\d*(?:\.\d+)?(?:[eE][+-]?\d+)?)([A-Za-zµ]?)", unit) + if input_match is None: + return 1 + + input_unit = input_match.group(2) + if not input_unit: + return 1 + + return CvtTools._scale_for_prefix(input_unit[0]) + + @staticmethod + def parse_to_hz(freq: str, default_unit: str = "") -> float: + new_freq = CvtTools.parse_general_val(input=freq, default_unit=default_unit) + return float(new_freq) if new_freq else 0.0 + + @staticmethod + def parse_to_Vpp(vpp: str) -> float: + vpp = vpp.replace(" ", "") + vpp_match = re.search(r"([+-]?\d*(?:\.\d+)?(?:[eE][+-]?\d+)?)([A-Za-zµ]*)", vpp) + if vpp_match is None: + return 0.0 + + vpp_val = vpp_match.group(1) + vpp_unit = vpp_match.group(2) + + if not vpp_val: + return 0.0 + + value = float(vpp_val) + if not vpp_unit: + scale = 1.0 + elif "Vpp".lower() in vpp_unit.lower(): + scale = 1.0 + elif "Vpk".lower() in vpp_unit.lower(): + scale = 2.0 + elif "Vrms".lower() in vpp_unit.lower(): + scale = math.sqrt(8) + else: + scale = 1.0 + + if vpp_unit and vpp_unit[0] == "m": + scale *= 0.001 + return value * scale + + @staticmethod + def parse_to_V(volts: str) -> float | int: + return CvtTools.parse_general_val(input=volts) + + @staticmethod + def _parabolic_interp_delta(m1: float, m0: float, p1: float) -> float: + eps = 1e-30 + m1 = np.log(max(m1, eps)) + m0 = np.log(max(m0, eps)) + p1 = np.log(max(p1, eps)) + denom = m1 - 2 * m0 + p1 + if abs(denom) < 1e-12: + return 0.0 + return 0.5 * (m1 - p1) / denom + + @staticmethod + def _complex_tone_at(times: np.ndarray, volts_ac: np.ndarray, f_hz: float, window: np.ndarray | None = None) -> complex: + if window is None: + return np.sum(volts_ac * np.exp(-1j * 2 * np.pi * f_hz * times)) + return np.sum(window * volts_ac * np.exp(-1j * 2 * np.pi * f_hz * times)) + + @staticmethod + def _scale_for_prefix(prefix: str) -> float: + if prefix in ("G", "g"): + return 1e9 + if prefix == "M": + return 1e6 + if prefix in ("k", "K"): + return 1e3 + if prefix == "m": + return 1e-3 + if prefix in ("u", "µ", "μ"): + return 1e-6 + if prefix in ("n", "N"): + return 1e-9 + if prefix in ("p", "P"): + return 1e-12 + return 1.0 diff --git a/src/app/shared/mapping.py b/src/app/shared/mapping.py new file mode 100644 index 0000000..09247e2 --- /dev/null +++ b/src/app/shared/mapping.py @@ -0,0 +1,170 @@ +from __future__ import annotations + + +class Mapping: + label_for_input_ui = "Input Control Panel" + label_for_file_menu = "File" + label_for_config_menu = "Device Manager" + label_for_device_configure_window = "Advanced Settings" + label_for_exit = "Exit" + + label_for_auto_lan = "Connection Mode" + label_for_visa_address = "VISA address" + label_for_ip_address = "IP address" + label_for_auto = "Auto" + label_for_lan = "LAN" + + label_for_chan_index = "Chan" + label_for_test_chan = "Test Chan" + label_for_ref_chan = "Ref Chan" + label_for_trig_chan = "Trig Chan" + + label_for_set_start_frequency = "Start Freq" + label_for_set_stop_frequency = "Stop Freq" + label_for_set_step_freq = "Freq Step" + label_for_set_step_num = "Step Ct" + label_for_set_center_frequency = "Center Freq" + label_for_set_interval_frequency = "Freq Span" + label_for_log = "Log" + label_for_freq_unit = "Unit" + label_for_points = "Max Samp" + label_for_freq = "Freq" + label_for_set_amp = "Amp" + label_for_set_imp = "Imp" + label_for_imp_r50 = "R50" + label_for_imp_inf = "High-Z" + label_for_coup = "Coup" + label_for_yoffset = "Center V" + label_for_range = "Full-Scale V" + label_for_auto_range = "Auto" + + label_for_single_chan_correct = "Single Chan" + label_for_duo_chan_correct = "Dual Chan" + label_for_no_correct = "No Cali" + label_for_set_ref = "Set As Ref" + label_for_load_ref = "Load Ref" + label_for_enable_ref = "Enable Cali" + + label_for_figure_gain = "Gain" + label_for_figure_gain_db = "dB" + label_for_figure_phase = "Phase (deg)" + label_for_figure_gain_freq = f"{label_for_figure_gain}_vs_{label_for_freq}" + label_for_figure_gaindb_freq = f"{label_for_figure_gain_db}_vs_{label_for_freq}" + + label_for_load_file_to_show = "Load for display" + label_for_load_file_to_ref = "Load for ref" + label_for_load_config = "Load config" + label_for_save_file = "Save data" + label_for_save_config = "Save config" + label_for_file_is_saved = "File saved" + label_for_sub_folder_data = "__data__" + + error_file_not_save = "Failed to save data!!!" + error_fail_auto_save = "Automatic save failed!!!" + title_alert = "Warning" + + mapping_auto_detect = "Auto Detect" + label_for_device_type_awg = "AWG" + label_for_device_type_osc = "OSC" + + mapping_DSG_4102 = "DSG4102" + mapping_DSG_836 = "DSG836" + mapping_MDO_34 = "MDO34" + mapping_MDO_3024 = "MDO3024" + mapping_DHO_1202 = "DHO1202" + mapping_DHO_1204 = "DHO1204" + + mapping_hz = "Hz" + mapping_khz = "KHz" + mapping_mhz = "MHz" + mapping_ghz = "GHz" + + mapping_imp_r50 = "50" + mapping_imp_high_z = "INF" + + mapping_vpp = "Vpp" + mapping_vpk = "Vpk" + mapping_vrms = "Vrms" + + mapping_file_ext_mat = ".mat" + mapping_file_ext_csv = ".csv" + mapping_file_ext_txt = ".txt" + mapping_file_ext_png = ".png" + + mapping_coup_ac = "AC" + mapping_coup_dc = "DC" + + mapping_state_on = "ON" + mapping_state_off = "OFF" + + values_awg = [mapping_DSG_4102, mapping_DSG_836] + values_osc = [mapping_MDO_34, mapping_MDO_3024, mapping_DHO_1202, mapping_DHO_1204] + values_device_type = [label_for_device_type_awg, label_for_device_type_osc] + values_freq_unit = [mapping_hz, mapping_khz, mapping_mhz, mapping_ghz] + values_device_num_list = [1, 2, 3, 4] + values_test_load_off_figure = [label_for_figure_gain_freq, label_for_figure_gaindb_freq] + values_correct_modes = [label_for_no_correct, label_for_single_chan_correct, label_for_duo_chan_correct] + values_coup = [mapping_coup_ac, mapping_coup_dc] + + mapping_freq = "freq" + mapping_gain_raw = "gain_raw" + mapping_gain_db_raw = "gain_db_raw" + mapping_phase_deg = "phase" + mapping_gain_corr = "gain_corr" + mapping_gain_db_corr = "gain_db_corr" + mapping_phase_deg_corr = "phase_corr" + mapping_gain_complex = "gain_complex" + + label_for_free_run = "free run" + label_for_triggered = "triggered" + values_trig_mode = [label_for_free_run, label_for_triggered] + + mapping_color_for_phase_line = "tab:red" + + label_for_mag = "Magnitude" + label_for_phase = "Phase" + label_for_mag_and_phase = "Magnitude + Phase" + values_mag_or_phase = [label_for_mag, label_for_phase, label_for_mag_and_phase] + + default_data_fn = "Test_File" + default_show_selection_font = ("Microsoft YaHei", 18) + default_text_font = ("Microsoft YaHei", 10) + default_terminal_bg = "black" + default_terminal_fg = "white" + + default_start_freq = "1.0" + default_stop_freq = "100.0" + default_step_freq = "1.0" + default_step_num = "100" + default_is_log_freq_enabled = mapping_state_off + default_freq_unit = mapping_mhz + default_samp_pts = "10000" + + default_awg_amp = "1.0" + default_awg_imp = "50" + default_yoffset = "0.0" + default_range = "1.0" + default_osc_imp = "50" + default_osc_coup = mapping_coup_dc + + default_is_auto_range = mapping_state_on + default_correct_mode = label_for_no_correct + default_is_correct_enabled = mapping_state_off + default_trig_mode = label_for_free_run + default_is_auto_save = mapping_state_on + default_is_auto_reset = mapping_state_on + + default_awg_name = mapping_DSG_4102 + default_osc_name = mapping_MDO_34 + + default_awg_chan_index = "1" + default_osc_test_chan_index = "1" + default_osc_trig_chan_index = "2" + default_osc_ref_chan_index = "2" + + default_awg_connect_mode = label_for_auto + default_osc_connect_mode = label_for_auto + default_awg_visa = "" + default_osc_visa = "" + default_awg_ip = "0.0.0.0" + default_osc_ip = "0.0.0.0" diff --git a/src/cvtTools.py b/src/cvtTools.py index 2d3264a..c3e7df2 100644 --- a/src/cvtTools.py +++ b/src/cvtTools.py @@ -1,129 +1,3 @@ -import re -import math -import numpy as np - -class CvtTools: - - @staticmethod - def parse_general_val(input: str, default_unit: str=None) -> float|int: - """ - Parse a numeric string with an optional prefix and return the scaled value. - Empty or invalid text resolves to 0, and unknown prefixes default to a scale of 1. - """ - input = input.replace(" ", "") - input_match = re.search(r"([+-]?\d*(?:\.\d+)?(?:[eE][+-]?\d+)?)([A-Za-zµ]?)", input) - input_val = input_match.group(1) - input_unit = input_match.group(2) - - try: - input_val = int(input_val) if input_val.isdigit() else float(input_val) - except: - return 0.0 - - if not input_unit: - val = input_val * CvtTools.convert_general_unit(default_unit) - # Prefer ints for integral results so downstream callers can treat counts as integers. - try: - return int(val) if float(val).is_integer() else val - except Exception: - return val - prefix = input_unit[0] - - if prefix in ('G', 'g'): scale = 1e9 - elif prefix == 'M': scale = 1e6 - elif prefix in ('k', 'K'): scale = 1e3 - elif prefix == 'm': scale = 1e-3 - elif prefix in ('u', 'µ', 'μ'):scale = 1e-6 - elif prefix in ('n', 'N'): scale = 1e-9 - elif prefix in ('p', 'P'): scale = 1e-12 - else: scale = 1 - - val = input_val * scale - try: - return int(val) if float(val).is_integer() else val - except Exception: - return val - - @staticmethod - def convert_general_unit(unit: str) -> float|int: - """ - Parse the prefix multiplier only; empty strings and unknown prefixes return 1. - """ - if not unit: return 1 - input_match = re.search(r"([+-]?\d*(?:\.\d+)?(?:[eE][+-]?\d+)?)([A-Za-zµ]?)", unit) - input_unit = input_match.group(2) - - if not input_unit: return 1 - prefix = input_unit[0] - - if prefix in ('G', 'g'): scale = 1e9 - elif prefix == 'M': scale = 1e6 - elif prefix in ('k', 'K'): scale = 1e3 - elif prefix == 'm': scale = 1e-3 - elif prefix in ('u', 'µ', 'μ'):scale = 1e-6 - elif prefix in ('n', 'N'): scale = 1e-9 - elif prefix in ('p', 'P'): scale = 1e-12 - else: scale = 1 - - return scale - - @staticmethod - def parse_to_hz(freq: str, default_unit: str = "") -> float: - """ - Parse frequency text with an optional default unit. - """ - new_freq = CvtTools.parse_general_val(input=freq, default_unit=default_unit) - - return new_freq if new_freq else 0 - - @staticmethod - def parse_to_Vpp(vpp: str) -> float: - """ - Parse voltage text and convert it to Vpp. - Supports: Vpp = 1x, Vpk = 2x, Vrms = sqrt(8) * Vpp, and an optional milli prefix. - """ - vpp = vpp.replace(" ", "") - vpp_macth = re.search(r"([+-]?\d*(?:\.\d+)?(?:[eE][+-]?\d+)?)([A-Za-zµ]*)", vpp) - vpp_val = vpp_macth.group(1) - vpp_unit = vpp_macth.group(2) - - if not vpp_val: return "" - vpp_val = float(vpp_val) - - if not vpp_unit: scale = 1 - elif "Vpp".lower() in vpp_unit.lower(): scale = 1 - elif "Vpk".lower() in vpp_unit.lower(): scale = 2 - elif "Vrms".lower() in vpp_unit.lower(): scale = math.sqrt(8) - else: scale = 1 - - if vpp_unit and vpp_unit[0] == "m": scale *= 0.001 - return vpp_val * scale - - @staticmethod - def parse_to_V(volts: str): - """ - Generic voltage parser helper. - """ - return CvtTools.parse_general_val(input=volts) - - @staticmethod - def _parabolic_interp_delta(m1, m0, p1): - """ - Estimate the peak offset via log|X| parabolic interpolation on three points. - """ - eps = 1e-30 - m1 = np.log(max(m1, eps)) - m0 = np.log(max(m0, eps)) - p1 = np.log(max(p1, eps)) - denom = (m1 - 2*m0 + p1) - if abs(denom) < 1e-12: return 0.0 - return 0.5 * (m1 - p1) / denom - - @staticmethod - def _complex_tone_at(times, volts_ac, f_hz, window=None): - """ - Compute a single-point DFT at f_hz and return the complex coefficient. - """ - if window is None: - return np.sum(volts_ac * np.exp(-1j * 2*np.pi * f_hz * times)) - return np.sum(window * volts_ac * np.exp(-1j * 2*np.pi * f_hz * times)) +from app.shared.cvt_tools import CvtTools + +__all__ = ["CvtTools"] diff --git a/src/equips.py b/src/equips.py index b59c67e..49b0773 100644 --- a/src/equips.py +++ b/src/equips.py @@ -1,14 +1,10 @@ -# Environment setup notes -# cd C:\Users\15038\Desktop\HardWare\mm_report -# python -m venv venv -# .\venv\Scripts\Activate.ps1 -# python issues.py -# pip install requests PyGithub pyinstaller -# pyinstaller --onefile --name mm_test.exe equips_v0.py - -"""equip driver Module deal with base equips operation. - new equips are encouraged to be inherited from bATEinst_base -""" +"""Legacy vendor/instrument compatibility layer. + +The refactored application does not call this module from UI or use cases. +It is wrapped by infrastructure adapters under ``app.infrastructure.instruments``. +Keep hardware-command changes conservative unless they can be verified on real +instruments. +""" import time import traceback diff --git a/src/main.py b/src/main.py index ed7953a..0c9cacc 100644 --- a/src/main.py +++ b/src/main.py @@ -1,37 +1,10 @@ from __future__ import annotations -from app.application.use_cases.load_measurement import LoadMeasurementUseCase -from app.application.use_cases.load_reference import LoadReferenceUseCase -from app.application.use_cases.save_measurement import SaveMeasurementUseCase -from app.application.use_cases.settings_use_case import SettingsUseCase -from app.infrastructure.instruments.resource_scanner import PyVisaResourceScanner -from app.infrastructure.persistence.measurement_repo_mat_csv import MatCsvMeasurementRepository -from app.infrastructure.persistence.reference_repo_mat import MatReferenceRepository -from app.infrastructure.persistence.settings_repo_json import JsonSettingsRepository -from app.presentation.tk.app_window import AppWindow -from app.presentation.tk.controller import TkController - def main() -> None: - window = AppWindow() - vm = window.vm - - settings_repo = JsonSettingsRepository() - measurement_repo = MatCsvMeasurementRepository() - reference_repo = MatReferenceRepository() - - controller = TkController( - window=window, - vm=vm, - settings_use_case=SettingsUseCase(settings_repo), - save_measurement_use_case=SaveMeasurementUseCase(measurement_repo), - load_measurement_use_case=LoadMeasurementUseCase(measurement_repo), - load_reference_use_case=LoadReferenceUseCase(reference_repo), - scanner=PyVisaResourceScanner(), - ) - controller.initialize() + from app.bootstrap import run_desktop_app - window.mainloop() + run_desktop_app() if __name__ == "__main__": diff --git a/src/mapping.py b/src/mapping.py index 40a0340..7bc0b48 100644 --- a/src/mapping.py +++ b/src/mapping.py @@ -1,216 +1,3 @@ -from __future__ import annotations +from app.shared.mapping import Mapping -class Mapping: - # ========================= messages ========================= - label_for_input_ui = "Input Control Panel" - label_for_file_menu = "File" - label_for_config_menu = "Device Manager" - label_for_device_configure_window = "Advanced Settings" - label_for_exit = "Exit" - - # ========================= connection ========================= - label_for_auto_lan = "Connection Mode" - label_for_visa_address = "VISA address" - label_for_ip_address = "IP address" - label_for_auto = "Auto" - label_for_lan = "LAN" - - # ========================= channel ========================= - label_for_chan_index = "Chan" - label_for_test_chan = "Test Chan" - label_for_ref_chan = "Ref Chan" - label_for_trig_chan = "Trig Chan" - - # ========================= awg/osc setting ========================= - label_for_set_start_frequency = "Start Freq" - label_for_set_stop_frequency = "Stop Freq" - label_for_set_step_freq = "Freq Step" - label_for_set_step_num = "Step Ct" - label_for_set_center_frequency = "Center Freq" - label_for_set_interval_frequency = "Freq Span" - label_for_log = "Log" - label_for_freq_unit = "Unit" - label_for_points = "Max Samp" - label_for_freq = "Freq" - label_for_set_amp = "Amp" - label_for_set_imp = "Imp" - label_for_imp_r50 = "R50" - label_for_imp_inf = "High-Z" - label_for_coup = "Coup" - label_for_yoffset = "Center V" - label_for_range = "Full-Scale V" - label_for_auto_range = "Auto" - - # ========================= correct/reference ========================= - label_for_single_chan_correct = "Single Chan" - label_for_duo_chan_correct = "Dual Chan" - label_for_no_correct = "No Cali" - label_for_set_ref = "Set As Ref" - label_for_load_ref = "Load Ref" - label_for_enable_ref = "Enable Cali" - - # ========================= figure ========================= - label_for_figure_gain = "Gain" - label_for_figure_gain_db = "dB" - label_for_figure_phase = "Phase (deg)" - label_for_figure_gain_freq = f"{label_for_figure_gain}_vs_{label_for_freq}" - label_for_figure_gaindb_freq = f"{label_for_figure_gain_db}_vs_{label_for_freq}" - - # ========================= load/save data ========================= - label_for_load_file_to_show = "Load for display" - label_for_load_file_to_ref = "Load for ref" - label_for_load_config = "Load config" - label_for_save_file = "Save data" - label_for_save_config = "Save config" - label_for_file_is_saved = "File saved" - label_for_sub_folder_data = "__data__" - - # ========================= errors & titles ========================= - error_file_not_save = "Failed to save data!!!" - error_fail_auto_save = "Automatic save failed!!!" - title_alert = "Warning" - - # ========================= mappings / options ========================= - mapping_auto_detect = "Auto Detect" - label_for_device_type_awg = "AWG" - label_for_device_type_osc = "OSC" - - mapping_DSG_4102 = "DSG4102" - mapping_DSG_836 = "DSG836" - mapping_MDO_34 = "MDO34" - mapping_MDO_3024 = "MDO3024" - mapping_DHO_1202 = "DHO1202" - mapping_DHO_1204 = "DHO1204" - - mapping_hz = "Hz" - mapping_khz = "KHz" - mapping_mhz = "MHz" - mapping_ghz = "GHz" - - mapping_imp_r50 = "50" - mapping_imp_high_z = "INF" - - mapping_vpp = "Vpp" - mapping_vpk = "Vpk" - mapping_vrms = "Vrms" - - mapping_file_ext_mat = ".mat" - mapping_file_ext_csv = ".csv" - mapping_file_ext_txt = ".txt" - mapping_file_ext_png = ".png" - - mapping_coup_ac = "AC" - mapping_coup_dc = "DC" - - mapping_state_on = "ON" - mapping_state_off = "OFF" - - # ========================= combobox values ========================= - values_awg = [ - mapping_DSG_4102, - mapping_DSG_836, - ] - values_osc = [ - mapping_MDO_34, - mapping_MDO_3024, - mapping_DHO_1202, - mapping_DHO_1204, - ] - values_device_type = [ - label_for_device_type_awg, - label_for_device_type_osc, - ] - values_freq_unit = [ - mapping_hz, - mapping_khz, - mapping_mhz, - mapping_ghz, - ] - values_device_num_list = [1, 2, 3, 4] - values_test_load_off_figure = [ - label_for_figure_gain_freq, - label_for_figure_gaindb_freq, - ] - values_correct_modes = [ - label_for_no_correct, - label_for_single_chan_correct, - label_for_duo_chan_correct, - ] - values_coup = [ - mapping_coup_ac, - mapping_coup_dc, - ] - - # ========================= data keys ========================= - mapping_freq = "freq" - mapping_gain_raw = "gain_raw" - mapping_gain_db_raw = "gain_db_raw" - mapping_phase_deg = "phase" - mapping_gain_corr = "gain_corr" - mapping_gain_db_corr = "gain_db_corr" - mapping_phase_deg_corr = "phase_corr" - mapping_gain_complex = "gain_complex" - - # ========================= trigger modes ========================= - label_for_free_run = "free run" - label_for_triggered = "triggered" - values_trig_mode = [label_for_free_run, label_for_triggered] - - # ========================= colors ========================= - mapping_color_for_phase_line = "tab:red" - - # ========================= plot options ========================= - label_for_mag = "Magnitude" - label_for_phase = "Phase" - label_for_mag_and_phase = "Magnitude + Phase" - values_mag_or_phase = [label_for_mag, label_for_phase, label_for_mag_and_phase] - - # ========================= defaults (UI/fonts/theme) ========================= - default_data_fn = "Test_File" - default_show_selection_font = ("Microsoft YaHei", 18) - default_text_font = ("Microsoft YaHei", 10) - default_terminal_bg = "black" - default_terminal_fg = "white" - - # ========================= defaults (general sweep) ========================= - default_start_freq = "1.0" - default_stop_freq = "100.0" - default_step_freq = "1.0" - default_step_num = "100" - default_is_log_freq_enabled = mapping_state_off - default_freq_unit = mapping_mhz - default_samp_pts = "10000" - - # ========================= defaults (AWG/OSC settings) ========================= - default_awg_amp = "1.0" - default_awg_imp = "50" - default_yoffset = "0.0" - default_range = "1.0" - default_osc_imp = "50" - default_osc_coup = mapping_coup_dc - - # ========================= defaults (modes/switches) ========================= - default_is_auto_range = mapping_state_on - default_correct_mode = label_for_no_correct - default_is_correct_enabled = mapping_state_off - default_trig_mode = label_for_free_run - default_is_auto_save = mapping_state_on - default_is_auto_reset = mapping_state_on - - # ========================= defaults (device names) ========================= - default_awg_name = mapping_DSG_4102 - default_osc_name = mapping_MDO_34 - - # ========================= defaults (channel indices) ========================= - default_awg_chan_index = "1" - default_osc_test_chan_index = "1" - default_osc_trig_chan_index = "2" - default_osc_ref_chan_index = "2" - - # ========================= defaults (connection) ========================= - default_awg_connect_mode = label_for_auto - default_osc_connect_mode = label_for_auto - default_awg_visa = "" - default_osc_visa = "" - default_awg_ip = "0.0.0.0" - default_osc_ip = "0.0.0.0" +__all__ = ["Mapping"] diff --git a/tests/test_app_paths.py b/tests/test_app_paths.py new file mode 100644 index 0000000..18ecb4d --- /dev/null +++ b/tests/test_app_paths.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import os +import sys +from pathlib import Path +import tempfile +import unittest +from unittest.mock import patch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + +from app.runtime.paths import APP_ROOT_ENV, AppPaths + + +class AppPathsTests(unittest.TestCase): + def test_from_root_derives_runtime_paths(self) -> None: + with tempfile.TemporaryDirectory() as td: + paths = AppPaths.from_root(Path(td)) + + self.assertEqual(paths.config_dir.name, "__config__") + self.assertEqual(paths.data_dir.name, "__data__") + self.assertEqual(paths.settings_path.name, "settings.json") + self.assertEqual(paths.measurement_dir.name, "measurement") + + def test_default_uses_working_directory(self) -> None: + with tempfile.TemporaryDirectory() as td, patch.dict(os.environ, {}, clear=True): + previous = Path.cwd() + try: + os.chdir(td) + paths = AppPaths.default() + finally: + os.chdir(previous) + + self.assertEqual(paths.root_dir, Path(td).resolve()) + + def test_default_can_be_overridden_by_environment(self) -> None: + with tempfile.TemporaryDirectory() as td, patch.dict(os.environ, {APP_ROOT_ENV: td}, clear=True): + paths = AppPaths.default() + + self.assertEqual(paths.root_dir, Path(td).resolve()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_architecture_boundaries.py b/tests/test_architecture_boundaries.py new file mode 100644 index 0000000..2bf0710 --- /dev/null +++ b/tests/test_architecture_boundaries.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import ast +import sys +from pathlib import Path +import unittest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +SRC_APP = PROJECT_ROOT / "src" / "app" + + +def imported_modules(path: Path) -> set[str]: + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + modules: set[str] = set() + for node in ast.walk(tree): + if isinstance(node, ast.Import): + modules.update(alias.name for alias in node.names) + elif isinstance(node, ast.ImportFrom) and node.module: + modules.add(node.module) + return modules + + +def py_files(base: Path) -> list[Path]: + return [path for path in base.rglob("*.py") if "__pycache__" not in path.parts] + + +class ArchitectureBoundaryTests(unittest.TestCase): + def test_domain_stays_pure(self) -> None: + forbidden_prefixes = ("tkinter", "pyvisa", "serial", "matplotlib", "app.infrastructure", "app.presentation") + offenders = [] + for path in py_files(SRC_APP / "domain"): + for module in imported_modules(path): + if module.startswith(forbidden_prefixes): + offenders.append((path.relative_to(PROJECT_ROOT), module)) + + self.assertEqual(offenders, []) + + def test_application_does_not_import_infrastructure_or_presentation(self) -> None: + forbidden_prefixes = ("app.infrastructure", "app.presentation") + offenders = [] + for path in py_files(SRC_APP / "application"): + for module in imported_modules(path): + if module.startswith(forbidden_prefixes): + offenders.append((path.relative_to(PROJECT_ROOT), module)) + + self.assertEqual(offenders, []) + + def test_presentation_does_not_import_infrastructure(self) -> None: + offenders = [] + for path in py_files(SRC_APP / "presentation"): + for module in imported_modules(path): + if module.startswith("app.infrastructure"): + offenders.append((path.relative_to(PROJECT_ROOT), module)) + + self.assertEqual(offenders, []) + + +if __name__ == "__main__": + unittest.main() + diff --git a/tests/test_auto_range_policy.py b/tests/test_auto_range_policy.py new file mode 100644 index 0000000..ab4cb24 --- /dev/null +++ b/tests/test_auto_range_policy.py @@ -0,0 +1,54 @@ +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.auto_range import AutoRangePolicy + + +class AutoRangePolicyTests(unittest.TestCase): + def setUp(self) -> None: + self.policy = AutoRangePolicy() + + def test_expand_range_when_signal_near_clipping(self) -> None: + decision = self.policy.decide( + volts=np.array([-0.45, 0.45]), + current_range_v=1.0, + current_offset_v=0.0, + requested_offset_v=0.0, + ) + + self.assertTrue(decision.changed) + self.assertGreater(decision.target_range_v, 1.0) + self.assertAlmostEqual(decision.target_offset_v, 0.0) + + def test_shrink_range_when_signal_small(self) -> None: + decision = self.policy.decide( + volts=np.array([-0.1, 0.1]), + current_range_v=1.0, + current_offset_v=0.0, + requested_offset_v=0.0, + ) + + self.assertTrue(decision.changed) + self.assertAlmostEqual(decision.target_range_v, 0.5, delta=1e-6) + + def test_adjust_offset_when_midpoint_drifts(self) -> None: + decision = self.policy.decide( + volts=np.array([0.35, 0.45, 0.55]), + current_range_v=1.0, + current_offset_v=0.0, + requested_offset_v=0.0, + ) + + self.assertTrue(decision.changed) + self.assertAlmostEqual(decision.target_offset_v, 0.45, delta=1e-6) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_measurement_io.py b/tests/test_measurement_io.py new file mode 100644 index 0000000..7c59b28 --- /dev/null +++ b/tests/test_measurement_io.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +import math +import sys +from pathlib import Path +import tempfile +import unittest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + +from app.application.dto import SaveTarget +from app.domain.models import SweepPoint, SweepResult +from app.infrastructure.persistence.measurement_exporter import MeasurementExporter +from app.infrastructure.persistence.measurement_loader import MeasurementLoader +from app.infrastructure.persistence.settings_defaults import DefaultSettingsFactory + + +class MeasurementIOTests(unittest.TestCase): + def test_export_and_load_round_trip_for_mat_and_csv(self) -> None: + settings = DefaultSettingsFactory().create() + result = SweepResult( + points=[ + SweepPoint(freq_hz=1_000.0, gain_linear=1.0, gain_db=0.0, phase_deg=5.0), + SweepPoint(freq_hz=2_000.0, gain_linear=2.0, gain_db=20.0 * math.log10(2.0), phase_deg=10.0), + ] + ) + exporter = MeasurementExporter() + loader = MeasurementLoader() + + with tempfile.TemporaryDirectory() as td: + artifacts = exporter.export( + result=result, + settings=settings, + target=SaveTarget(base_path=Path(td) / "measurement", include_timestamp=False, figures={}), + ) + + loaded_mat = loader.load(str(artifacts.mat_path)) + loaded_csv = loader.load(str(artifacts.csv_path)) + + self.assertEqual(len(loaded_mat.result.points), 2) + self.assertEqual(len(loaded_csv.result.points), 2) + self.assertAlmostEqual(loaded_mat.result.points[1].gain_linear, 2.0, delta=1e-6) + self.assertAlmostEqual(loaded_csv.result.points[1].phase_deg or 0.0, 10.0, delta=1e-6) + + def test_sparse_phase_values_keep_row_alignment(self) -> None: + settings = DefaultSettingsFactory().create() + result = SweepResult( + points=[ + SweepPoint(freq_hz=1_000.0, gain_linear=1.0, gain_db=0.0, phase_deg=None), + SweepPoint(freq_hz=2_000.0, gain_linear=2.0, gain_db=20.0 * math.log10(2.0), phase_deg=10.0), + ] + ) + exporter = MeasurementExporter() + loader = MeasurementLoader() + + with tempfile.TemporaryDirectory() as td: + artifacts = exporter.export( + result=result, + settings=settings, + target=SaveTarget(base_path=Path(td) / "measurement", include_timestamp=False, figures={}), + ) + loaded_csv = loader.load(str(artifacts.csv_path)) + csv_text = artifacts.csv_path.read_text(encoding="utf-8") + + self.assertIsNone(loaded_csv.result.points[0].phase_deg) + self.assertAlmostEqual(loaded_csv.result.points[1].phase_deg or 0.0, 10.0, delta=1e-6) + self.assertIn("1000.0,1.0,0.0,", csv_text) + + def test_demo_data_mat_files_load(self) -> None: + loader = MeasurementLoader() + demo_dir = Path(__file__).resolve().parents[1] / "demo_data" + + raw = loader.load(str(demo_dir / "Deme(1).mat")) + corrected = loader.load(str(demo_dir / "Demo(2).mat")) + + self.assertEqual(len(raw.result.points), 15) + self.assertEqual(len(corrected.result.points), 50) + self.assertIsNone(raw.result.points[0].phase_deg) + self.assertIsNotNone(corrected.result.points[0].phase_deg) + self.assertGreater(raw.result.points[0].gain_linear, 0.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_point_measurement_services.py b/tests/test_point_measurement_services.py new file mode 100644 index 0000000..89c5ab6 --- /dev/null +++ b/tests/test_point_measurement_services.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +import math +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.application.dto import StartSweepCommand +from app.application.services.sweep.calibration_applier import CalibrationApplier +from app.application.services.sweep.models import AcquiredPointData +from app.application.services.sweep.point_measurement_service import PointMeasurementService +from app.domain.enums import ConnectionMode, CorrectionMode, CouplingMode, ImpedanceMode, MagnitudePhaseMode, TriggerMode +from app.domain.models import ( + AppSettings, + AwgSettings, + ChannelSelection, + InstrumentEndpoint, + InstrumentSetup, + OscSettings, + RunMode, + SweepPoint, + SweepSpec, +) + + +def build_settings(*, correction_mode: CorrectionMode, trigger_mode: TriggerMode) -> AppSettings: + return AppSettings( + schema_version=1, + freq_unit="Hz", + sweep=SweepSpec(start_hz=1_000.0, stop_hz=1_000.0, step_hz=1_000.0, step_count=None, is_log=False), + run_mode=RunMode( + correction_mode=correction_mode, + trigger_mode=trigger_mode, + auto_range=False, + auto_reset=True, + ), + setup=InstrumentSetup( + awg=InstrumentEndpoint(model="DSG4102", connect_mode=ConnectionMode.AUTO), + osc=InstrumentEndpoint(model="MDO34", connect_mode=ConnectionMode.AUTO), + channels=ChannelSelection(awg_ch=1, osc_test_ch=1, osc_ref_ch=2, osc_trig_ch=2), + awg_settings=AwgSettings(amplitude_vpp=1.0, impedance=ImpedanceMode.R50), + osc_settings=OscSettings( + full_scale_v=1.0, + offset_v=0.0, + points=4000, + impedance=ImpedanceMode.R50, + coupling=CouplingMode.DC, + ), + ), + magnitude_phase_mode=MagnitudePhaseMode.MAG, + auto_save_data=False, + ) + + +class PointMeasurementServiceTests(unittest.TestCase): + def setUp(self) -> None: + self.service = PointMeasurementService() + self.calibration = CalibrationApplier() + + def test_single_channel_triggered_returns_phase(self) -> None: + fs = 200_000 + f0 = 2_000 + t = np.arange(0.0, 0.03, 1.0 / fs) + volts = 0.5 * np.sin(2.0 * np.pi * f0 * t) + + point = self.service.measure( + settings=build_settings(correction_mode=CorrectionMode.NONE, trigger_mode=TriggerMode.TRIGGERED), + acquired=AcquiredPointData( + actual_freq_hz=f0, + read_amp_vpp=1.0, + test_times=t, + test_volts=volts, + ), + ) + + self.assertAlmostEqual(point.gain_linear, 1.0, delta=0.05) + self.assertIsNotNone(point.phase_deg) + self.assertIsNotNone(point.gain_complex) + + def test_dual_channel_measurement_returns_ratio(self) -> None: + fs = 200_000 + f0 = 5_000 + t = np.arange(0.0, 0.03, 1.0 / fs) + ref = np.sin(2.0 * np.pi * f0 * t) + test = 2.0 * np.sin(2.0 * np.pi * f0 * t + math.radians(30.0)) + + point = self.service.measure( + settings=build_settings(correction_mode=CorrectionMode.DUAL, trigger_mode=TriggerMode.FREE_RUN), + acquired=AcquiredPointData( + actual_freq_hz=f0, + read_amp_vpp=1.0, + test_times=t, + test_volts=test, + ref_times=t, + ref_volts=ref, + ), + ) + + self.assertAlmostEqual(point.gain_linear, 2.0, delta=0.1) + self.assertAlmostEqual(point.phase_deg or 0.0, 30.0, delta=5.0) + + def test_calibration_applier_corrects_gain_and_phase(self) -> None: + point = SweepPoint( + freq_hz=5_000.0, + gain_linear=2.0, + gain_db=20.0 * math.log10(2.0), + phase_deg=30.0, + gain_complex=2.0 * np.exp(1j * np.deg2rad(30.0)), + ) + cmd = StartSweepCommand( + settings=build_settings(correction_mode=CorrectionMode.DUAL, trigger_mode=TriggerMode.TRIGGERED), + calibration_enabled=True, + reference_interpolator=lambda _xs: np.array([2.0 * np.exp(1j * np.deg2rad(10.0))]), + ) + + corrected = self.calibration.apply(point=point, cmd=cmd) + self.assertAlmostEqual(corrected.gain_linear, 1.0, delta=1e-6) + self.assertAlmostEqual(corrected.phase_deg or 0.0, 20.0, delta=1e-6) + + def test_calibration_applier_handles_zero_reference(self) -> None: + point = SweepPoint( + freq_hz=5_000.0, + gain_linear=1.0, + gain_db=0.0, + phase_deg=0.0, + gain_complex=1.0 + 0.0j, + ) + cmd = StartSweepCommand( + settings=build_settings(correction_mode=CorrectionMode.DUAL, trigger_mode=TriggerMode.TRIGGERED), + calibration_enabled=True, + reference_interpolator=lambda _xs: np.array([0.0 + 0.0j]), + ) + + corrected = self.calibration.apply(point=point, cmd=cmd) + self.assertTrue(math.isfinite(corrected.gain_linear)) + self.assertTrue(math.isfinite(corrected.gain_db)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_settings_serializer.py b/tests/test_settings_serializer.py new file mode 100644 index 0000000..82088ff --- /dev/null +++ b/tests/test_settings_serializer.py @@ -0,0 +1,30 @@ +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 TriggerMode +from app.infrastructure.persistence.settings_defaults import DefaultSettingsFactory +from app.infrastructure.persistence.settings_serializer import SettingsSerializer + + +class SettingsSerializerTests(unittest.TestCase): + def test_schema_version_one_round_trip(self) -> None: + serializer = SettingsSerializer() + settings = DefaultSettingsFactory().create() + settings.run_mode.trigger_mode = TriggerMode.TRIGGERED + settings.setup.awg.visa_address = "USB::MOCK::INSTR" + + payload = serializer.to_payload(settings) + loaded = serializer.from_payload(payload) + + self.assertEqual(loaded.schema_version, 1) + self.assertEqual(loaded.run_mode.trigger_mode, TriggerMode.TRIGGERED) + self.assertEqual(loaded.setup.awg.visa_address, "USB::MOCK::INSTR") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_start_sweep_use_case.py b/tests/test_start_sweep_use_case.py index 1b21ac5..a03b1b8 100644 --- a/tests/test_start_sweep_use_case.py +++ b/tests/test_start_sweep_use_case.py @@ -43,6 +43,9 @@ def reset(self) -> None: def output_on(self, channel: int) -> None: _ = channel + def output_off(self, channel: int) -> None: + _ = channel + def set_impedance(self, mode: str, channel: int) -> None: _ = (mode, channel) diff --git a/tests/test_sweep_planner_service.py b/tests/test_sweep_planner_service.py new file mode 100644 index 0000000..e81efcc --- /dev/null +++ b/tests/test_sweep_planner_service.py @@ -0,0 +1,79 @@ +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.application.services.sweep.planner import SweepPlanner +from app.domain.enums import ConnectionMode, CorrectionMode, CouplingMode, ImpedanceMode, MagnitudePhaseMode, TriggerMode +from app.domain.models import ( + AppSettings, + AwgSettings, + ChannelSelection, + InstrumentEndpoint, + InstrumentSetup, + OscSettings, + RunMode, + SweepSpec, +) + + +def build_settings(*, is_log: bool) -> AppSettings: + return AppSettings( + schema_version=1, + freq_unit="Hz", + sweep=SweepSpec( + start_hz=1.0, + stop_hz=100.0 if is_log else 5.0, + step_hz=None if is_log else 2.0, + step_count=5 if is_log else None, + is_log=is_log, + ), + run_mode=RunMode( + correction_mode=CorrectionMode.NONE, + trigger_mode=TriggerMode.FREE_RUN, + auto_range=False, + auto_reset=True, + ), + setup=InstrumentSetup( + awg=InstrumentEndpoint(model="DSG4102", connect_mode=ConnectionMode.AUTO), + osc=InstrumentEndpoint(model="MDO34", connect_mode=ConnectionMode.AUTO), + channels=ChannelSelection(awg_ch=1, osc_test_ch=1, osc_ref_ch=2, osc_trig_ch=2), + awg_settings=AwgSettings(amplitude_vpp=1.0, impedance=ImpedanceMode.R50), + osc_settings=OscSettings( + full_scale_v=1.0, + offset_v=0.0, + points=10_000, + impedance=ImpedanceMode.R50, + coupling=CouplingMode.DC, + ), + ), + magnitude_phase_mode=MagnitudePhaseMode.MAG, + auto_save_data=False, + ) + + +class SweepPlannerServiceTests(unittest.TestCase): + def setUp(self) -> None: + self.planner = SweepPlanner() + + def test_linear_plan_matches_expected_points(self) -> None: + plan = self.planner.plan(build_settings(is_log=False)) + self.assertEqual(plan.freq_points.tolist(), [1.0, 3.0, 5.0]) + self.assertEqual(plan.total_points, 3) + + def test_log_plan_uses_step_count(self) -> None: + plan = self.planner.plan(build_settings(is_log=True)) + self.assertEqual(plan.total_points, 5) + self.assertAlmostEqual(float(plan.freq_points[0]), 1.0) + self.assertAlmostEqual(float(plan.freq_points[-1]), 100.0) + + def test_sampling_window_is_positive(self) -> None: + window = self.planner.compute_sampling_window_s(freq_hz=1e3, sample_rate_hz=1e6, points=10_000) + self.assertGreater(window, 0.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_sweep_task_runner.py b/tests/test_sweep_task_runner.py new file mode 100644 index 0000000..6926d7a --- /dev/null +++ b/tests/test_sweep_task_runner.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +import sys +from pathlib import Path +import threading +import unittest +from types import SimpleNamespace + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + +from app.application.events import SweepCompleted, 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 + + +class FakeSaveMeasurementUseCase: + def __init__(self) -> None: + self.calls: list[tuple[object, object, object]] = [] + + def execute(self, result, settings, target): + self.calls.append((result, settings, target)) + return SimpleNamespace(mat_path=Path("measurement.mat")) + + +class FakeEmitter: + def __init__(self) -> None: + self.events: list[object] = [] + + def emit(self, event: object) -> None: + self.events.append(event) + + +class FakePort: + def __init__(self) -> None: + self.closed = False + self.output_off_channels: list[int] = [] + self.fail_close = False + + def output_off(self, channel: int) -> None: + self.output_off_channels.append(channel) + + def close(self) -> None: + if self.fail_close: + raise RuntimeError("close failed") + self.closed = True + + +class FakeStartUseCase: + def __init__(self, result: SweepResult) -> None: + self._result = result + + def run(self, cmd, emitter): + emitter.emit(SweepCompleted(result=self._result)) + return self._result + + +class SweepTaskRunnerTests(unittest.TestCase): + def test_runner_auto_saves_completed_result(self) -> None: + settings = DefaultSettingsFactory().create() + settings.auto_save_data = True + save_use_case = FakeSaveMeasurementUseCase() + 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=save_use_case, + 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.assertEqual(len(save_use_case.calls), 1) + self.assertEqual(awg.output_off_channels, [settings.setup.channels.awg_ch]) + self.assertTrue(awg.closed) + self.assertTrue(osc.closed) + self.assertTrue(any(isinstance(event, SweepCompleted) for event in emitter.events)) + + def test_runner_emits_warning_when_port_close_fails(self) -> None: + settings = DefaultSettingsFactory().create() + save_use_case = FakeSaveMeasurementUseCase() + emitter = FakeEmitter() + awg = FakePort() + osc = FakePort() + awg.fail_close = True + 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=save_use_case, + 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) + + warnings = [event for event in emitter.events if isinstance(event, SweepWarning)] + self.assertTrue(any(event.code == "AWG_CLOSE_FAILED" for event in warnings)) + + +if __name__ == "__main__": + unittest.main()