Files
HVideoTool/scripts/smoke_test.py
T

248 lines
9.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Headless smoke tests for HVideoTool.
No formal test suite and no model weights in the repo, so this exercises the pure
core logic (detection cache, cross-model NMS, list-filter predicate, ETA formatting,
extract-dialog options, per-project settings round-trip) plus a minimal offscreen
``MainWindow`` build on a tiny throwaway project. It avoids torch/ultralytics entirely
(detection results are injected directly into ``_results``).
Run::
set QT_QPA_PLATFORM=offscreen
set PYTHONIOENCODING=utf-8
.venv\\Scripts\\python.exe scripts\\smoke_test.py
Exits non-zero on the first failure.
"""
from __future__ import annotations
import os
import sys
import tempfile
from pathlib import Path
import numpy as np
# Run Qt without a display and keep Unicode console output sane on Windows.
os.environ.setdefault("QT_QPA_PLATFORM", "offscreen")
os.environ.setdefault("PYTHONIOENCODING", "utf-8")
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from hvideotool.config import AppConfig # noqa: E402
from hvideotool.core.detection import cache as detection_cache # noqa: E402
from hvideotool.core.detection.multi import MultiYoloDetector, _iou, _nms # noqa: E402
from hvideotool.core.detection.types import CensorType, Detection # noqa: E402
from hvideotool.core.imageio import imwrite_unicode # noqa: E402
from hvideotool.core.project import Project # noqa: E402
_failures: list[str] = []
def check(cond: bool, msg: str) -> None:
status = "PASS" if cond else "FAIL"
print(f" [{status}] {msg}")
if not cond:
_failures.append(msg)
def _det(score: float, bbox, *, label="", model="", type=CensorType.MOSAIC) -> Detection:
return Detection(type=type, score=score, bbox=bbox, label=label, model=model)
# --------------------------------------------------------------------- core tests
def test_cache_atomic_roundtrip() -> None:
print("cache: atomic save/load round-trip")
with tempfile.TemporaryDirectory() as d:
base = Path(d)
cache_file = base / "detections.json"
key = detection_cache.make_key(["m/a.pt", "m/b.pt"], 0.2, 640)
results = {
str(base / "001.jpg"): [_det(0.9, (1, 2, 3, 4), label="mosaic", model="a")],
str(base / "002.jpg"): [], # checked-clean
}
ok = detection_cache.save_results(cache_file, key, results)
check(ok, "save_results returns True")
check(cache_file.is_file(), "cache file created")
check(not (base / "detections.json.tmp").exists(), "temp file cleaned up")
loaded = detection_cache.load_results(cache_file, key, base)
check(loaded is not None, "load returns a dict for matching key")
check(set(loaded) == set(results), "all basenames round-trip")
d1 = loaded[str(base / "001.jpg")][0]
check(d1.score == 0.9 and d1.model == "a", "detection fields preserved")
check(loaded[str(base / "002.jpg")] == [], "empty (clean) entry preserved")
bad = detection_cache.load_results(
cache_file, detection_cache.make_key(["x.pt"], 0.2, 640), base
)
check(bad is None, "mismatched detector key => None (cache ignored)")
# NMS only enters the key when enabled: default key stays valid, NMS key differs.
k_off = detection_cache.make_key(["a.pt"], 0.2, 640)
k_off2 = detection_cache.make_key(["a.pt"], 0.2, 640, nms_iou=None)
k_on = detection_cache.make_key(["a.pt"], 0.2, 640, nms_iou=0.6)
check(k_off == k_off2 and "nms_iou" not in k_off, "NMS-off key unchanged (no nms field)")
check(k_on != k_off and k_on.get("nms_iou") == 0.6, "NMS-on key is distinct")
def test_nms() -> None:
print("detection: cross-model NMS")
check(abs(_iou((0, 0, 10, 10), (0, 0, 10, 10)) - 1.0) < 1e-9, "IoU identical = 1.0")
check(_iou((0, 0, 10, 10), (100, 100, 5, 5)) == 0.0, "IoU disjoint = 0.0")
# Two near-duplicate boxes + one distinct: NMS keeps the higher score + the distinct.
dets = [
_det(0.6, (0, 0, 10, 10), model="a"),
_det(0.9, (1, 1, 10, 10), model="b"), # overlaps the first heavily
_det(0.8, (200, 200, 10, 10), model="c"), # separate region
]
kept = _nms(dets, 0.5)
check(len(kept) == 2, "two duplicates merged to one (+ the distinct box)")
check(any(k.score == 0.9 for k in kept), "higher-score duplicate survives")
check(not any(k.score == 0.6 for k in kept), "lower-score duplicate dropped")
class _Stub:
def __init__(self, ds):
self._ds = ds
@property
def name(self):
return "stub"
def detect(self, frame):
return list(self._ds)
merged = MultiYoloDetector(
[_Stub([dets[0]]), _Stub([dets[1], dets[2]])], nms_iou=0.5
).detect(None)
check(len(merged) == 2, "MultiYoloDetector applies NMS across detectors")
no_nms = MultiYoloDetector([_Stub([dets[0]]), _Stub([dets[1], dets[2]])]).detect(None)
check(len(no_nms) == 3, "without nms_iou, all detections are concatenated")
def test_extract_dialog_options() -> None:
print("extract dialog: options() includes JPEG quality")
from PySide6.QtWidgets import QApplication
from hvideotool.ui.extract_dialog import ExtractDialog
_ensure_app(QApplication)
dlg = ExtractDialog()
opts = dlg.options()
check(len(opts) == 4, "options() is a 4-tuple (keyframes, step, max_dim, quality)")
keyframes, step, max_dim, quality = opts
check(keyframes is False and step == 1, "defaults to every-frame (step=1)")
check(1 <= quality <= 100, "quality in 1..100")
def test_settings_roundtrip() -> None:
print("project: per-project settings round-trip (new fields)")
with tempfile.TemporaryDirectory() as d:
cfg = AppConfig()
cfg.model_thresholds = {"penis": 0.42}
cfg.cross_model_nms = True
cfg.nms_iou = 0.55
cfg.default_threshold = 0.3
proj = Project.create(Path(d) / "P", name="P")
proj.update_from_config(cfg)
proj.save()
reloaded = Project.load(proj.root)
cfg2 = AppConfig()
reloaded.apply_to_config(cfg2)
check(cfg2.model_thresholds == {"penis": 0.42}, "model_thresholds persisted")
check(cfg2.cross_model_nms is True, "cross_model_nms persisted")
check(abs(cfg2.nms_iou - 0.55) < 1e-9, "nms_iou persisted")
check(not (proj.root / "project.json.tmp").exists(), "project.json.tmp cleaned up")
# ---------------------------------------------------------------- GUI smoke tests
_APP = None
def _ensure_app(QApplication):
global _APP
if _APP is None:
_APP = QApplication.instance() or QApplication([])
return _APP
def test_mainwindow_filter_and_jump() -> None:
print("MainWindow: build, filter, jump, ETA (offscreen)")
from PySide6.QtWidgets import QApplication
from hvideotool.ui.main_window import MainWindow
_ensure_app(QApplication)
with tempfile.TemporaryDirectory() as d:
proj = Project.create(Path(d) / "Proj", name="Proj")
img = np.zeros((16, 16, 3), dtype=np.uint8)
names = [f"{i:03d}.jpg" for i in range(1, 6)]
for n in names:
imwrite_unicode(str(proj.frames_dir / n), img)
w = MainWindow(AppConfig())
w._open_project(proj)
check(w.file_list.count() == 5, "all 5 frames listed")
files = list(w._files)
# Inject detection results: frame0 has a hit, frame1 is clean, rest uncomputed.
w._results[str(files[0])] = [_det(0.9, (1, 1, 4, 4), label="mosaic", model="a")]
w._results[str(files[1])] = []
w._relabel_all()
w._refresh_marks()
# Filter: "С цензурой" (hits) shows only frame0.
w.filter_combo.setCurrentIndex(w.filter_combo.findData("hits"))
vis = [i for i in range(w.file_list.count()) if not w.file_list.item(i).isHidden()]
check(vis == [0], "filter 'hits' shows only the censored frame")
# Filter: "Не рассчитано" shows the 3 uncomputed frames.
w.filter_combo.setCurrentIndex(w.filter_combo.findData("uncomputed"))
vis = [i for i in range(w.file_list.count()) if not w.file_list.item(i).isHidden()]
check(vis == [2, 3, 4], "filter 'uncomputed' shows the not-yet-detected frames")
# _step skips hidden rows: from row2, next visible is row3.
w.file_list.setCurrentRow(2)
w._step(1)
check(w.file_list.currentRow() == 3, "_step skips filtered-out rows")
# Back to all.
w.filter_combo.setCurrentIndex(w.filter_combo.findData("all"))
vis = [i for i in range(w.file_list.count()) if not w.file_list.item(i).isHidden()]
check(len(vis) == 5, "filter 'all' shows everything again")
# Jump-to-frame core (bypassing the modal dialog): selecting row 4.
w.file_list.setCurrentRow(4)
check(w.pos_label.text() == "5 / 5", "position label reflects the current frame")
# ETA formatting.
w._job_start = __import__("time").monotonic() - 10.0 # 10s elapsed
suffix = w._eta_suffix(2, 10) # 2/10 done in 10s => ~40s remaining
check("осталось" in suffix, "ETA suffix produced for an in-progress job")
check(w._eta_suffix(0, 10) == "" and w._eta_suffix(10, 10) == "",
"no ETA at 0% or 100%")
w.close()
def main() -> int:
tests = [
test_cache_atomic_roundtrip,
test_nms,
test_extract_dialog_options,
test_settings_roundtrip,
test_mainwindow_filter_and_jump,
]
for t in tests:
t()
print()
if _failures:
print(f"FAILED ({len(_failures)}):")
for f in _failures:
print(" -", f)
return 1
print("ALL SMOKE TESTS PASSED")
return 0
if __name__ == "__main__":
raise SystemExit(main())